From 045aeb2de5547a3fb6444ec314a80c226b476978 Mon Sep 17 00:00:00 2001 From: Evan Jarrett Date: Thu, 1 Jan 2026 21:19:38 -0600 Subject: [PATCH] add index table to mst so listRecords is more efficient --- cmd/hold/main.go | 13 +- pkg/appview/jetstream/backfill.go | 7 +- pkg/hold/pds/records.go | 251 ++++++++++++ pkg/hold/pds/records_test.go | 627 ++++++++++++++++++++++++++++++ pkg/hold/pds/server.go | 99 ++++- pkg/hold/pds/server_test.go | 328 ++++++++++++++++ pkg/hold/pds/status_test.go | 2 +- pkg/hold/pds/xrpc.go | 166 ++++++-- pkg/hold/pds/xrpc_test.go | 296 ++++++++++++++ 9 files changed, 1748 insertions(+), 41 deletions(-) create mode 100644 pkg/hold/pds/records.go create mode 100644 pkg/hold/pds/records_test.go diff --git a/cmd/hold/main.go b/cmd/hold/main.go index 1d14030..f9c6420 100644 --- a/cmd/hold/main.go +++ b/cmd/hold/main.go @@ -82,10 +82,17 @@ func main() { slog.Warn("Failed to bootstrap events from repo", "error", err) } - // Wire up repo event handler to broadcaster - holdPDS.RepomgrRef().SetEventHandler(broadcaster.SetRepoEventHandler(), true) + // Backfill records index from existing MST data (one-time on startup) + if err := holdPDS.BackfillRecordsIndex(ctx); err != nil { + slog.Warn("Failed to backfill records index", "error", err) + } - slog.Info("Embedded PDS initialized successfully with firehose enabled") + // Wire up repo event handler with records indexing + broadcaster + // The indexing handler wraps the broadcaster handler to keep index in sync + indexingHandler := holdPDS.CreateRecordsIndexEventHandler(broadcaster.SetRepoEventHandler()) + holdPDS.RepomgrRef().SetEventHandler(indexingHandler, true) + + slog.Info("Embedded PDS initialized successfully with firehose and records index enabled") } else { slog.Error("Database path is required for embedded PDS authorization") os.Exit(1) diff --git a/pkg/appview/jetstream/backfill.go b/pkg/appview/jetstream/backfill.go index a44517a..6980ff9 100644 --- a/pkg/appview/jetstream/backfill.go +++ b/pkg/appview/jetstream/backfill.go @@ -50,7 +50,7 @@ func NewBackfillWorker(database *sql.DB, relayEndpoint, defaultHoldDID string, t return &BackfillWorker{ db: database, client: client, // This points to the relay - processor: NewProcessor(database, false, nil), // No cache for batch processing, no stats + processor: NewProcessor(database, false, NewStatsCache()), // Stats cache for aggregation defaultHoldDID: defaultHoldDID, testMode: testMode, refresher: refresher, @@ -76,6 +76,7 @@ func (b *BackfillWorker) Start(ctx context.Context) error { atproto.StarCollection, // io.atcr.sailor.star atproto.SailorProfileCollection, // io.atcr.sailor.profile atproto.RepoPageCollection, // io.atcr.repo.page + atproto.StatsCollection, // io.atcr.hold.stats (from holds) } for _, collection := range collections { @@ -311,6 +312,10 @@ func (b *BackfillWorker) processRecord(ctx context.Context, did, collection stri case atproto.RepoPageCollection: // rkey is extracted from the record URI, but for repo pages we use Repository field return b.processor.ProcessRepoPage(ctx, did, record.URI, record.Value, false) + case atproto.StatsCollection: + // Stats are stored in hold PDSes, not user PDSes + // 'did' here is the hold's DID (e.g., did:web:hold01.atcr.io) + return b.processor.ProcessStats(ctx, did, record.Value, false) default: return fmt.Errorf("unsupported collection: %s", collection) } diff --git a/pkg/hold/pds/records.go b/pkg/hold/pds/records.go new file mode 100644 index 0000000..5d1191c --- /dev/null +++ b/pkg/hold/pds/records.go @@ -0,0 +1,251 @@ +package pds + +import ( + "context" + "database/sql" + "fmt" + "log/slog" + "strings" + + "github.com/bluesky-social/indigo/repo" + "github.com/ipfs/go-cid" + _ "github.com/mattn/go-sqlite3" +) + +// RecordsIndex provides an efficient index for listing records +// This follows the official ATProto PDS pattern of using SQL for queries +// while MST is used for sync operations. +type RecordsIndex struct { + db *sql.DB +} + +// Record represents a record in the index +type Record struct { + Collection string + Rkey string + Cid string +} + +const recordsSchema = ` +CREATE TABLE IF NOT EXISTS records ( + collection TEXT NOT NULL, + rkey TEXT NOT NULL, + cid TEXT NOT NULL, + PRIMARY KEY (collection, rkey) +); +CREATE INDEX IF NOT EXISTS idx_records_collection_rkey ON records(collection, rkey); +` + +// NewRecordsIndex creates or opens a records index +func NewRecordsIndex(dbPath string) (*RecordsIndex, error) { + db, err := sql.Open("sqlite3", dbPath) + if err != nil { + return nil, fmt.Errorf("failed to open records database: %w", err) + } + + // Create schema + _, err = db.Exec(recordsSchema) + if err != nil { + db.Close() + return nil, fmt.Errorf("failed to create records schema: %w", err) + } + + return &RecordsIndex{db: db}, nil +} + +// Close closes the database connection +func (ri *RecordsIndex) Close() error { + if ri.db != nil { + return ri.db.Close() + } + return nil +} + +// IndexRecord adds or updates a record in the index +func (ri *RecordsIndex) IndexRecord(collection, rkey, cidStr string) error { + _, err := ri.db.Exec(` + INSERT OR REPLACE INTO records (collection, rkey, cid) + VALUES (?, ?, ?) + `, collection, rkey, cidStr) + return err +} + +// DeleteRecord removes a record from the index +func (ri *RecordsIndex) DeleteRecord(collection, rkey string) error { + _, err := ri.db.Exec(` + DELETE FROM records WHERE collection = ? AND rkey = ? + `, collection, rkey) + return err +} + +// ListRecords returns records for a collection with pagination support +// reverse=false (default): newest first (rkey DESC) +// reverse=true: oldest first (rkey ASC) +func (ri *RecordsIndex) ListRecords(collection string, limit int, cursor string, reverse bool) ([]Record, string, error) { + // Build query based on sort order + var query string + var args []any + + if reverse { + // Oldest first (ascending order) + if cursor != "" { + query = ` + SELECT collection, rkey, cid FROM records + WHERE collection = ? AND rkey > ? + ORDER BY rkey ASC + LIMIT ? + ` + args = []any{collection, cursor, limit + 1} + } else { + query = ` + SELECT collection, rkey, cid FROM records + WHERE collection = ? + ORDER BY rkey ASC + LIMIT ? + ` + args = []any{collection, limit + 1} + } + } else { + // Newest first (descending order) - default + if cursor != "" { + query = ` + SELECT collection, rkey, cid FROM records + WHERE collection = ? AND rkey < ? + ORDER BY rkey DESC + LIMIT ? + ` + args = []any{collection, cursor, limit + 1} + } else { + query = ` + SELECT collection, rkey, cid FROM records + WHERE collection = ? + ORDER BY rkey DESC + LIMIT ? + ` + args = []any{collection, limit + 1} + } + } + + rows, err := ri.db.Query(query, args...) + if err != nil { + return nil, "", fmt.Errorf("failed to query records: %w", err) + } + defer rows.Close() + + var records []Record + for rows.Next() { + var rec Record + if err := rows.Scan(&rec.Collection, &rec.Rkey, &rec.Cid); err != nil { + return nil, "", fmt.Errorf("failed to scan record: %w", err) + } + records = append(records, rec) + } + + if err := rows.Err(); err != nil { + return nil, "", fmt.Errorf("error iterating records: %w", err) + } + + // Determine next cursor + var nextCursor string + if len(records) > limit { + // More records available, set cursor to the last included record + nextCursor = records[limit-1].Rkey + records = records[:limit] + } + + return records, nextCursor, nil +} + +// Count returns the number of records in a collection +func (ri *RecordsIndex) Count(collection string) (int, error) { + var count int + err := ri.db.QueryRow(` + SELECT COUNT(*) FROM records WHERE collection = ? + `, collection).Scan(&count) + return count, err +} + +// TotalCount returns the total number of records in the index +func (ri *RecordsIndex) TotalCount() (int, error) { + var count int + err := ri.db.QueryRow(`SELECT COUNT(*) FROM records`).Scan(&count) + return count, err +} + +// BackfillFromRepo populates the records index from an existing MST repo +// Compares MST count with index count - only backfills if they differ +func (ri *RecordsIndex) BackfillFromRepo(ctx context.Context, repoHandle *repo.Repo) error { + // Count records in MST + mstCount := 0 + err := repoHandle.ForEach(ctx, "", func(key string, c cid.Cid) error { + mstCount++ + return nil + }) + if err != nil { + return fmt.Errorf("failed to count MST records: %w", err) + } + + // Count records in index + indexCount, err := ri.TotalCount() + if err != nil { + return fmt.Errorf("failed to check index count: %w", err) + } + + // Skip if counts match + if indexCount == mstCount { + slog.Debug("Records index in sync with MST", "count", indexCount) + return nil + } + + slog.Info("Backfilling records index from MST...", "mstCount", mstCount, "indexCount", indexCount) + + // Begin transaction for bulk insert + tx, err := ri.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("failed to begin transaction: %w", err) + } + defer tx.Rollback() + + stmt, err := tx.Prepare(` + INSERT OR REPLACE INTO records (collection, rkey, cid) + VALUES (?, ?, ?) + `) + if err != nil { + return fmt.Errorf("failed to prepare statement: %w", err) + } + defer stmt.Close() + + recordCount := 0 + err = repoHandle.ForEach(ctx, "", func(key string, c cid.Cid) error { + // key format: "collection/rkey" + parts := strings.SplitN(key, "/", 2) + if len(parts) != 2 { + return nil // Skip malformed keys + } + collection, rkey := parts[0], parts[1] + + _, err := stmt.Exec(collection, rkey, c.String()) + if err != nil { + return fmt.Errorf("failed to index record %s: %w", key, err) + } + recordCount++ + + // Log progress every 1000 records + if recordCount%1000 == 0 { + slog.Debug("Backfill progress", "count", recordCount) + } + + return nil + }) + + if err != nil { + return fmt.Errorf("failed to walk repo: %w", err) + } + + if err := tx.Commit(); err != nil { + return fmt.Errorf("failed to commit transaction: %w", err) + } + + slog.Info("Backfill complete", "records", recordCount) + return nil +} diff --git a/pkg/hold/pds/records_test.go b/pkg/hold/pds/records_test.go new file mode 100644 index 0000000..6b5b8e1 --- /dev/null +++ b/pkg/hold/pds/records_test.go @@ -0,0 +1,627 @@ +package pds + +import ( + "context" + "os" + "path/filepath" + "testing" + + "github.com/bluesky-social/indigo/repo" + _ "github.com/mattn/go-sqlite3" +) + +// Tests for RecordsIndex + +// TestNewRecordsIndex tests creating a new records index +func TestNewRecordsIndex(t *testing.T) { + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "records.db") + + ri, err := NewRecordsIndex(dbPath) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + if ri.db == nil { + t.Error("Expected db to be non-nil") + } + + // Verify database file was created + if _, err := os.Stat(dbPath); os.IsNotExist(err) { + t.Error("Expected database file to be created") + } +} + +// TestNewRecordsIndex_InvalidPath tests error handling for invalid path +func TestNewRecordsIndex_InvalidPath(t *testing.T) { + // Try to create in a non-existent directory + _, err := NewRecordsIndex("/nonexistent/dir/records.db") + if err == nil { + t.Error("Expected error for invalid path") + } +} + +// TestRecordsIndex_IndexRecord tests adding records to the index +func TestRecordsIndex_IndexRecord(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Index a record + err = ri.IndexRecord("io.atcr.hold.crew", "abc123", "bafyrei123") + if err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + + // Verify it was indexed + count, err := ri.Count("io.atcr.hold.crew") + if err != nil { + t.Fatalf("Count() error = %v", err) + } + if count != 1 { + t.Errorf("Expected count 1, got %d", count) + } +} + +// TestRecordsIndex_IndexRecord_Upsert tests updating an existing record +func TestRecordsIndex_IndexRecord_Upsert(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Index a record + err = ri.IndexRecord("io.atcr.hold.crew", "abc123", "bafyrei123") + if err != nil { + t.Fatalf("IndexRecord() first call error = %v", err) + } + + // Update the same record with new CID + err = ri.IndexRecord("io.atcr.hold.crew", "abc123", "bafyrei456") + if err != nil { + t.Fatalf("IndexRecord() second call error = %v", err) + } + + // Count should still be 1 (upsert, not insert) + count, err := ri.Count("io.atcr.hold.crew") + if err != nil { + t.Fatalf("Count() error = %v", err) + } + if count != 1 { + t.Errorf("Expected count 1 after upsert, got %d", count) + } + + // Verify the CID was updated + records, _, err := ri.ListRecords("io.atcr.hold.crew", 10, "", false) + if err != nil { + t.Fatalf("ListRecords() error = %v", err) + } + if len(records) != 1 { + t.Fatalf("Expected 1 record, got %d", len(records)) + } + if records[0].Cid != "bafyrei456" { + t.Errorf("Expected CID bafyrei456, got %s", records[0].Cid) + } +} + +// TestRecordsIndex_DeleteRecord tests removing a record from the index +func TestRecordsIndex_DeleteRecord(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Index a record + err = ri.IndexRecord("io.atcr.hold.crew", "abc123", "bafyrei123") + if err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + + // Delete it + err = ri.DeleteRecord("io.atcr.hold.crew", "abc123") + if err != nil { + t.Fatalf("DeleteRecord() error = %v", err) + } + + // Verify it was deleted + count, err := ri.Count("io.atcr.hold.crew") + if err != nil { + t.Fatalf("Count() error = %v", err) + } + if count != 0 { + t.Errorf("Expected count 0 after delete, got %d", count) + } +} + +// TestRecordsIndex_DeleteRecord_NotExists tests deleting a non-existent record +func TestRecordsIndex_DeleteRecord_NotExists(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Delete a record that doesn't exist - should not error + err = ri.DeleteRecord("io.atcr.hold.crew", "nonexistent") + if err != nil { + t.Errorf("DeleteRecord() should not error for non-existent record, got: %v", err) + } +} + +// TestRecordsIndex_Close tests clean shutdown +func TestRecordsIndex_Close(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + + err = ri.Close() + if err != nil { + t.Errorf("Close() error = %v", err) + } + + // Double close should not panic (nil check) + ri.db = nil + err = ri.Close() + if err != nil { + t.Errorf("Close() on nil db error = %v", err) + } +} + +// TestRecordsIndex_ListRecords_Empty tests listing an empty collection +func TestRecordsIndex_ListRecords_Empty(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + records, cursor, err := ri.ListRecords("io.atcr.hold.crew", 10, "", false) + if err != nil { + t.Fatalf("ListRecords() error = %v", err) + } + + if len(records) != 0 { + t.Errorf("Expected empty records, got %d", len(records)) + } + if cursor != "" { + t.Errorf("Expected empty cursor, got %s", cursor) + } +} + +// TestRecordsIndex_ListRecords_Basic tests basic listing +func TestRecordsIndex_ListRecords_Basic(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add some records + records := []struct { + rkey string + cid string + }{ + {"aaa", "cid1"}, + {"bbb", "cid2"}, + {"ccc", "cid3"}, + } + for _, r := range records { + if err := ri.IndexRecord("io.atcr.hold.crew", r.rkey, r.cid); err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + } + + // List all + result, cursor, err := ri.ListRecords("io.atcr.hold.crew", 10, "", false) + if err != nil { + t.Fatalf("ListRecords() error = %v", err) + } + + if len(result) != 3 { + t.Errorf("Expected 3 records, got %d", len(result)) + } + if cursor != "" { + t.Errorf("Expected no cursor when all records returned, got %s", cursor) + } +} + +// TestRecordsIndex_ListRecords_DefaultOrder tests newest-first ordering (DESC) +func TestRecordsIndex_ListRecords_DefaultOrder(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add records with different rkeys (TIDs are lexicographically ordered by time) + rkeys := []string{"3m3aaaaaaaaa", "3m3bbbbbbbbb", "3m3ccccccccc"} + for _, rkey := range rkeys { + if err := ri.IndexRecord("io.atcr.hold.crew", rkey, "cid-"+rkey); err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + } + + // List with default order (newest first = DESC) + records, _, err := ri.ListRecords("io.atcr.hold.crew", 10, "", false) + if err != nil { + t.Fatalf("ListRecords() error = %v", err) + } + + // Should be in descending order + if len(records) != 3 { + t.Fatalf("Expected 3 records, got %d", len(records)) + } + if records[0].Rkey != "3m3ccccccccc" { + t.Errorf("Expected first record rkey=3m3ccccccccc, got %s", records[0].Rkey) + } + if records[1].Rkey != "3m3bbbbbbbbb" { + t.Errorf("Expected second record rkey=3m3bbbbbbbbb, got %s", records[1].Rkey) + } + if records[2].Rkey != "3m3aaaaaaaaa" { + t.Errorf("Expected third record rkey=3m3aaaaaaaaa, got %s", records[2].Rkey) + } +} + +// TestRecordsIndex_ListRecords_ReverseOrder tests oldest-first ordering (ASC) +func TestRecordsIndex_ListRecords_ReverseOrder(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add records + rkeys := []string{"3m3aaaaaaaaa", "3m3bbbbbbbbb", "3m3ccccccccc"} + for _, rkey := range rkeys { + if err := ri.IndexRecord("io.atcr.hold.crew", rkey, "cid-"+rkey); err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + } + + // List with reverse=true (oldest first = ASC) + records, _, err := ri.ListRecords("io.atcr.hold.crew", 10, "", true) + if err != nil { + t.Fatalf("ListRecords() error = %v", err) + } + + // Should be in ascending order + if len(records) != 3 { + t.Fatalf("Expected 3 records, got %d", len(records)) + } + if records[0].Rkey != "3m3aaaaaaaaa" { + t.Errorf("Expected first record rkey=3m3aaaaaaaaa, got %s", records[0].Rkey) + } + if records[1].Rkey != "3m3bbbbbbbbb" { + t.Errorf("Expected second record rkey=3m3bbbbbbbbb, got %s", records[1].Rkey) + } + if records[2].Rkey != "3m3ccccccccc" { + t.Errorf("Expected third record rkey=3m3ccccccccc, got %s", records[2].Rkey) + } +} + +// TestRecordsIndex_ListRecords_Limit tests the limit parameter +func TestRecordsIndex_ListRecords_Limit(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add 5 records + for i := 0; i < 5; i++ { + rkey := string(rune('a' + i)) + if err := ri.IndexRecord("io.atcr.hold.crew", rkey, "cid-"+rkey); err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + } + + // List with limit=2 + records, cursor, err := ri.ListRecords("io.atcr.hold.crew", 2, "", false) + if err != nil { + t.Fatalf("ListRecords() error = %v", err) + } + + if len(records) != 2 { + t.Errorf("Expected 2 records with limit=2, got %d", len(records)) + } + if cursor == "" { + t.Error("Expected cursor when more records exist") + } +} + +// TestRecordsIndex_ListRecords_Cursor tests pagination with cursor +func TestRecordsIndex_ListRecords_Cursor(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add 5 records + rkeys := []string{"a", "b", "c", "d", "e"} + for _, rkey := range rkeys { + if err := ri.IndexRecord("io.atcr.hold.crew", rkey, "cid-"+rkey); err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + } + + // First page (default order = DESC, so e, d first) + page1, cursor1, err := ri.ListRecords("io.atcr.hold.crew", 2, "", false) + if err != nil { + t.Fatalf("ListRecords() page 1 error = %v", err) + } + if len(page1) != 2 { + t.Fatalf("Expected 2 records in page 1, got %d", len(page1)) + } + if cursor1 == "" { + t.Fatal("Expected cursor after page 1") + } + + // Second page using cursor + page2, cursor2, err := ri.ListRecords("io.atcr.hold.crew", 2, cursor1, false) + if err != nil { + t.Fatalf("ListRecords() page 2 error = %v", err) + } + if len(page2) != 2 { + t.Errorf("Expected 2 records in page 2, got %d", len(page2)) + } + + // Third page + page3, cursor3, err := ri.ListRecords("io.atcr.hold.crew", 2, cursor2, false) + if err != nil { + t.Fatalf("ListRecords() page 3 error = %v", err) + } + if len(page3) != 1 { + t.Errorf("Expected 1 record in page 3, got %d", len(page3)) + } + if cursor3 != "" { + t.Errorf("Expected no cursor after last page, got %s", cursor3) + } + + // Verify no duplicates across pages + seen := make(map[string]bool) + for _, r := range page1 { + if seen[r.Rkey] { + t.Errorf("Duplicate record: %s", r.Rkey) + } + seen[r.Rkey] = true + } + for _, r := range page2 { + if seen[r.Rkey] { + t.Errorf("Duplicate record: %s", r.Rkey) + } + seen[r.Rkey] = true + } + for _, r := range page3 { + if seen[r.Rkey] { + t.Errorf("Duplicate record: %s", r.Rkey) + } + seen[r.Rkey] = true + } + if len(seen) != 5 { + t.Errorf("Expected 5 unique records, got %d", len(seen)) + } +} + +// TestRecordsIndex_ListRecords_CursorReverse tests pagination with reverse order +func TestRecordsIndex_ListRecords_CursorReverse(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add 5 records + rkeys := []string{"a", "b", "c", "d", "e"} + for _, rkey := range rkeys { + if err := ri.IndexRecord("io.atcr.hold.crew", rkey, "cid-"+rkey); err != nil { + t.Fatalf("IndexRecord() error = %v", err) + } + } + + // First page (reverse = ASC, so a, b first) + page1, cursor1, err := ri.ListRecords("io.atcr.hold.crew", 2, "", true) + if err != nil { + t.Fatalf("ListRecords() page 1 error = %v", err) + } + if len(page1) != 2 { + t.Fatalf("Expected 2 records in page 1, got %d", len(page1)) + } + if page1[0].Rkey != "a" { + t.Errorf("Expected first record a, got %s", page1[0].Rkey) + } + if page1[1].Rkey != "b" { + t.Errorf("Expected second record b, got %s", page1[1].Rkey) + } + + // Second page using cursor + page2, _, err := ri.ListRecords("io.atcr.hold.crew", 2, cursor1, true) + if err != nil { + t.Fatalf("ListRecords() page 2 error = %v", err) + } + if len(page2) != 2 { + t.Errorf("Expected 2 records in page 2, got %d", len(page2)) + } + if page2[0].Rkey != "c" { + t.Errorf("Expected first record c, got %s", page2[0].Rkey) + } +} + +// TestRecordsIndex_Count tests counting records in a collection +func TestRecordsIndex_Count(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add records to two collections + for i := 0; i < 3; i++ { + ri.IndexRecord("io.atcr.hold.crew", string(rune('a'+i)), "cid1") + } + for i := 0; i < 5; i++ { + ri.IndexRecord("io.atcr.hold.captain", string(rune('a'+i)), "cid2") + } + + // Count crew + count, err := ri.Count("io.atcr.hold.crew") + if err != nil { + t.Fatalf("Count() error = %v", err) + } + if count != 3 { + t.Errorf("Expected crew count 3, got %d", count) + } + + // Count captain + count, err = ri.Count("io.atcr.hold.captain") + if err != nil { + t.Fatalf("Count() error = %v", err) + } + if count != 5 { + t.Errorf("Expected captain count 5, got %d", count) + } +} + +// TestRecordsIndex_Count_Empty tests counting an empty collection +func TestRecordsIndex_Count_Empty(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + count, err := ri.Count("io.atcr.nonexistent") + if err != nil { + t.Fatalf("Count() error = %v", err) + } + if count != 0 { + t.Errorf("Expected count 0 for empty collection, got %d", count) + } +} + +// TestRecordsIndex_TotalCount tests total count across all collections +func TestRecordsIndex_TotalCount(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add records to multiple collections + ri.IndexRecord("io.atcr.hold.crew", "a", "cid1") + ri.IndexRecord("io.atcr.hold.crew", "b", "cid2") + ri.IndexRecord("io.atcr.hold.captain", "self", "cid3") + ri.IndexRecord("io.atcr.manifest", "abc123", "cid4") + + count, err := ri.TotalCount() + if err != nil { + t.Fatalf("TotalCount() error = %v", err) + } + if count != 4 { + t.Errorf("Expected total count 4, got %d", count) + } +} + +// TestRecordsIndex_BackfillFromRepo_Empty tests backfill with empty repo +func TestRecordsIndex_BackfillFromRepo_Empty(t *testing.T) { + // This test requires a mock repo which is complex to set up + // Skip for now - the integration tests in server_test.go will cover this + t.Skip("Requires mock repo setup - covered by integration tests") +} + +// TestRecordsIndex_BackfillFromRepo tests backfill from MST +func TestRecordsIndex_BackfillFromRepo(t *testing.T) { + // This test requires a real repo with MST data + // Skip unit test - covered by integration tests in server_test.go + t.Skip("Requires real repo with MST - covered by integration tests") +} + +// TestRecordsIndex_BackfillFromRepo_SkipsWhenSynced tests backfill skip logic +func TestRecordsIndex_BackfillFromRepo_SkipsWhenSynced(t *testing.T) { + // Create a mock scenario where counts match + // This is tested via the count comparison logic + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // The skip logic depends on count comparison in BackfillFromRepo + // which requires a real repo. Skip for now. + t.Skip("Requires mock repo - covered by integration tests") +} + +// TestRecordsIndex_MultipleCollections tests isolation between collections +func TestRecordsIndex_MultipleCollections(t *testing.T) { + tmpDir := t.TempDir() + ri, err := NewRecordsIndex(filepath.Join(tmpDir, "records.db")) + if err != nil { + t.Fatalf("NewRecordsIndex() error = %v", err) + } + defer ri.Close() + + // Add records to different collections with same rkeys + ri.IndexRecord("io.atcr.hold.crew", "abc", "cid-crew") + ri.IndexRecord("io.atcr.hold.captain", "abc", "cid-captain") + ri.IndexRecord("io.atcr.manifest", "abc", "cid-manifest") + + // Listing should only return records from requested collection + records, _, err := ri.ListRecords("io.atcr.hold.crew", 10, "", false) + if err != nil { + t.Fatalf("ListRecords() error = %v", err) + } + if len(records) != 1 { + t.Errorf("Expected 1 crew record, got %d", len(records)) + } + if records[0].Cid != "cid-crew" { + t.Errorf("Expected cid-crew, got %s", records[0].Cid) + } + + // Delete from one collection shouldn't affect others + ri.DeleteRecord("io.atcr.hold.crew", "abc") + + count, _ := ri.Count("io.atcr.hold.captain") + if count != 1 { + t.Errorf("Expected captain count 1 after deleting crew, got %d", count) + } +} + +// mockRepo is a minimal mock for testing backfill +// Note: Full backfill testing requires integration tests with real repo +type mockRepo struct { + records map[string]string // key -> cid +} + +func (m *mockRepo) ForEach(ctx context.Context, prefix string, fn func(string, interface{}) error) error { + for k, v := range m.records { + if err := fn(k, v); err != nil { + if err == repo.ErrDoneIterating { + return nil + } + return err + } + } + return nil +} diff --git a/pkg/hold/pds/server.go b/pkg/hold/pds/server.go index d6cfa68..1541de4 100644 --- a/pkg/hold/pds/server.go +++ b/pkg/hold/pds/server.go @@ -41,6 +41,7 @@ type HoldPDS struct { uid models.Uid signingKey *atcrypto.PrivateKeyK256 enableBlueskyPosts bool + recordsIndex *RecordsIndex } // NewHoldPDS creates or opens a hold PDS with SQLite carstore @@ -98,6 +99,17 @@ func NewHoldPDS(ctx context.Context, did, publicURL, dbPath, keyPath string, ena slog.Info("New hold repo - will be initialized in Bootstrap") } + // Initialize records index for efficient listing queries + // Uses same database as carstore for simplicity + var recordsIndex *RecordsIndex + if dbPath != ":memory:" { + recordsDbPath := dbPath + "/db.sqlite3" + recordsIndex, err = NewRecordsIndex(recordsDbPath) + if err != nil { + return nil, fmt.Errorf("failed to create records index: %w", err) + } + } + return &HoldPDS{ did: did, PublicURL: publicURL, @@ -107,6 +119,7 @@ func NewHoldPDS(ctx context.Context, did, publicURL, dbPath, keyPath string, ena uid: uid, signingKey: signingKey, enableBlueskyPosts: enableBlueskyPosts, + recordsIndex: recordsIndex, }, nil } @@ -125,6 +138,21 @@ func (p *HoldPDS) RepomgrRef() *RepoManager { return p.repomgr } +// RecordsIndex returns the records index for efficient listing +func (p *HoldPDS) RecordsIndex() *RecordsIndex { + return p.recordsIndex +} + +// Carstore returns the carstore for repo operations +func (p *HoldPDS) Carstore() carstore.CarStore { + return p.carstore +} + +// UID returns the user ID for this hold +func (p *HoldPDS) UID() models.Uid { + return p.uid +} + // Bootstrap initializes the hold with the captain record, owner as first crew member, and profile func (p *HoldPDS) Bootstrap(ctx context.Context, storageDriver driver.StorageDriver, ownerDID string, public bool, allowAllCrew bool, avatarURL string) error { if ownerDID == "" { @@ -268,8 +296,75 @@ func (p *HoldPDS) ListCollections(ctx context.Context) ([]string, error) { return result, nil } -// Close closes the carstore +// Close closes the carstore and records index func (p *HoldPDS) Close() error { - // TODO: Close session properly + if p.recordsIndex != nil { + if err := p.recordsIndex.Close(); err != nil { + return fmt.Errorf("failed to close records index: %w", err) + } + } return nil } + +// CreateRecordsIndexEventHandler creates an event handler that indexes records +// and also calls the provided broadcaster handler +func (p *HoldPDS) CreateRecordsIndexEventHandler(broadcasterHandler func(context.Context, *RepoEvent)) func(context.Context, *RepoEvent) { + return func(ctx context.Context, event *RepoEvent) { + // Index/delete records based on event operations + if p.recordsIndex != nil { + for _, op := range event.Ops { + switch op.Kind { + case EvtKindCreateRecord, EvtKindUpdateRecord: + // Index the record + cidStr := "" + if op.RecCid != nil { + cidStr = op.RecCid.String() + } + if err := p.recordsIndex.IndexRecord(op.Collection, op.Rkey, cidStr); err != nil { + slog.Warn("Failed to index record", "collection", op.Collection, "rkey", op.Rkey, "error", err) + } + case EvtKindDeleteRecord: + // Remove from index + if err := p.recordsIndex.DeleteRecord(op.Collection, op.Rkey); err != nil { + slog.Warn("Failed to delete record from index", "collection", op.Collection, "rkey", op.Rkey, "error", err) + } + } + } + } + + // Call the broadcaster handler + if broadcasterHandler != nil { + broadcasterHandler(ctx, event) + } + } +} + +// BackfillRecordsIndex populates the records index from existing MST data +func (p *HoldPDS) BackfillRecordsIndex(ctx context.Context) error { + if p.recordsIndex == nil { + return nil // No index to backfill + } + + // Create session to read repo + session, err := p.carstore.ReadOnlySession(p.uid) + if err != nil { + return fmt.Errorf("failed to create session: %w", err) + } + + head, err := p.carstore.GetUserRepoHead(ctx, p.uid) + if err != nil { + return fmt.Errorf("failed to get repo head: %w", err) + } + + if !head.Defined() { + slog.Debug("No repo head, skipping backfill") + return nil + } + + repoHandle, err := repo.OpenRepo(ctx, session, head) + if err != nil { + return fmt.Errorf("failed to open repo: %w", err) + } + + return p.recordsIndex.BackfillFromRepo(ctx, repoHandle) +} diff --git a/pkg/hold/pds/server_test.go b/pkg/hold/pds/server_test.go index d08c530..8d0f556 100644 --- a/pkg/hold/pds/server_test.go +++ b/pkg/hold/pds/server_test.go @@ -620,3 +620,331 @@ func TestBootstrap_CaptainWithoutCrew(t *testing.T) { } } } + +// Tests for RecordsIndex feature + +// TestHoldPDS_RecordsIndex_Nil tests that RecordsIndex is nil for :memory: database +func TestHoldPDS_RecordsIndex_Nil(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + keyPath := filepath.Join(tmpDir, "signing-key") + + // Create with :memory: database + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", ":memory:", keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // RecordsIndex should be nil for :memory: + if pds.RecordsIndex() != nil { + t.Error("Expected RecordsIndex() to be nil for :memory: database") + } +} + +// TestHoldPDS_RecordsIndex_NonNil tests that RecordsIndex is created for file database +func TestHoldPDS_RecordsIndex_NonNil(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + + // Create with file database + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", dbPath, keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // RecordsIndex should be non-nil for file database + if pds.RecordsIndex() == nil { + t.Error("Expected RecordsIndex() to be non-nil for file database") + } +} + +// TestHoldPDS_Carstore tests the Carstore getter +func TestHoldPDS_Carstore(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + keyPath := filepath.Join(tmpDir, "signing-key") + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", ":memory:", keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + if pds.Carstore() == nil { + t.Error("Expected Carstore() to be non-nil") + } +} + +// TestHoldPDS_UID tests the UID getter +func TestHoldPDS_UID(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + keyPath := filepath.Join(tmpDir, "signing-key") + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", ":memory:", keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // UID should be 1 for single-user PDS + if pds.UID() != 1 { + t.Errorf("Expected UID() to be 1, got %d", pds.UID()) + } +} + +// TestHoldPDS_CreateRecordsIndexEventHandler tests event handler wrapper +func TestHoldPDS_CreateRecordsIndexEventHandler(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", dbPath, keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // Track if broadcaster was called + broadcasterCalled := false + broadcasterHandler := func(ctx context.Context, event *RepoEvent) { + broadcasterCalled = true + } + + // Create handler + handler := pds.CreateRecordsIndexEventHandler(broadcasterHandler) + if handler == nil { + t.Fatal("Expected handler to be non-nil") + } + + // Create a test event with create operation + event := &RepoEvent{ + Ops: []RepoOp{ + { + Kind: EvtKindCreateRecord, + Collection: "io.atcr.hold.crew", + Rkey: "testrkey", + RecCid: nil, // Will be nil string + }, + }, + } + + // Call handler + handler(ctx, event) + + // Verify broadcaster was called + if !broadcasterCalled { + t.Error("Expected broadcaster handler to be called") + } + + // Verify record was indexed + if pds.RecordsIndex() != nil { + count, err := pds.RecordsIndex().Count("io.atcr.hold.crew") + if err != nil { + t.Fatalf("Count() error = %v", err) + } + if count != 1 { + t.Errorf("Expected 1 indexed record, got %d", count) + } + } +} + +// TestHoldPDS_CreateRecordsIndexEventHandler_Delete tests delete operation +func TestHoldPDS_CreateRecordsIndexEventHandler_Delete(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", dbPath, keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + handler := pds.CreateRecordsIndexEventHandler(nil) + + // First, create a record + createEvent := &RepoEvent{ + Ops: []RepoOp{ + { + Kind: EvtKindCreateRecord, + Collection: "io.atcr.hold.crew", + Rkey: "testrkey", + }, + }, + } + handler(ctx, createEvent) + + // Verify it was indexed + count, _ := pds.RecordsIndex().Count("io.atcr.hold.crew") + if count != 1 { + t.Fatalf("Expected 1 record after create, got %d", count) + } + + // Now delete it + deleteEvent := &RepoEvent{ + Ops: []RepoOp{ + { + Kind: EvtKindDeleteRecord, + Collection: "io.atcr.hold.crew", + Rkey: "testrkey", + }, + }, + } + handler(ctx, deleteEvent) + + // Verify it was removed from index + count, _ = pds.RecordsIndex().Count("io.atcr.hold.crew") + if count != 0 { + t.Errorf("Expected 0 records after delete, got %d", count) + } +} + +// TestHoldPDS_CreateRecordsIndexEventHandler_NilBroadcaster tests with nil broadcaster +func TestHoldPDS_CreateRecordsIndexEventHandler_NilBroadcaster(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", dbPath, keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // Create handler with nil broadcaster (should not panic) + handler := pds.CreateRecordsIndexEventHandler(nil) + + event := &RepoEvent{ + Ops: []RepoOp{ + { + Kind: EvtKindCreateRecord, + Collection: "io.atcr.hold.crew", + Rkey: "testrkey", + }, + }, + } + + // Should not panic + handler(ctx, event) + + // Verify record was still indexed + count, _ := pds.RecordsIndex().Count("io.atcr.hold.crew") + if count != 1 { + t.Errorf("Expected 1 indexed record, got %d", count) + } +} + +// TestHoldPDS_BackfillRecordsIndex tests backfilling the records index from MST +func TestHoldPDS_BackfillRecordsIndex(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", dbPath, keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // Bootstrap to create some records in MST (captain + crew) + ownerDID := "did:plc:testowner" + err = pds.Bootstrap(ctx, nil, ownerDID, true, false, "") + if err != nil { + t.Fatalf("Bootstrap failed: %v", err) + } + + // Clear the index to simulate out-of-sync state + _, err = pds.RecordsIndex().db.Exec("DELETE FROM records") + if err != nil { + t.Fatalf("Failed to clear index: %v", err) + } + + // Verify index is empty + count, _ := pds.RecordsIndex().TotalCount() + if count != 0 { + t.Fatalf("Expected empty index, got %d", count) + } + + // Backfill + err = pds.BackfillRecordsIndex(ctx) + if err != nil { + t.Fatalf("BackfillRecordsIndex failed: %v", err) + } + + // Verify records were backfilled + // Bootstrap creates: 1 captain + 1 crew + 1 profile = 3 records + count, _ = pds.RecordsIndex().TotalCount() + if count < 2 { + t.Errorf("Expected at least 2 records after backfill (captain + crew), got %d", count) + } +} + +// TestHoldPDS_BackfillRecordsIndex_NilIndex tests backfill with nil index +func TestHoldPDS_BackfillRecordsIndex_NilIndex(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + keyPath := filepath.Join(tmpDir, "signing-key") + + // Use :memory: to get nil index + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", ":memory:", keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // Backfill should be no-op and not error + err = pds.BackfillRecordsIndex(ctx) + if err != nil { + t.Errorf("BackfillRecordsIndex should not error with nil index, got: %v", err) + } +} + +// TestHoldPDS_BackfillRecordsIndex_SkipsWhenSynced tests backfill skip when already synced +func TestHoldPDS_BackfillRecordsIndex_SkipsWhenSynced(t *testing.T) { + ctx := context.Background() + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", dbPath, keyPath, false) + if err != nil { + t.Fatalf("NewHoldPDS failed: %v", err) + } + defer pds.Close() + + // Bootstrap to create records + err = pds.Bootstrap(ctx, nil, "did:plc:testowner", true, false, "") + if err != nil { + t.Fatalf("Bootstrap failed: %v", err) + } + + // Backfill once to sync + err = pds.BackfillRecordsIndex(ctx) + if err != nil { + t.Fatalf("First BackfillRecordsIndex failed: %v", err) + } + + count1, _ := pds.RecordsIndex().TotalCount() + + // Backfill again - should skip (counts match) + err = pds.BackfillRecordsIndex(ctx) + if err != nil { + t.Fatalf("Second BackfillRecordsIndex failed: %v", err) + } + + count2, _ := pds.RecordsIndex().TotalCount() + + // Count should be unchanged + if count1 != count2 { + t.Errorf("Expected count to remain %d after second backfill, got %d", count1, count2) + } +} diff --git a/pkg/hold/pds/status_test.go b/pkg/hold/pds/status_test.go index 430f36c..5b1b78b 100644 --- a/pkg/hold/pds/status_test.go +++ b/pkg/hold/pds/status_test.go @@ -63,7 +63,7 @@ func TestStatusPost(t *testing.T) { "repo": did, "collection": atproto.BskyPostCollection, "limit": "100", - "reverse": "true", // Most recent first + // Default order (reverse=false) is newest first (DESC by rkey) }) w := httptest.NewRecorder() handler.HandleListRecords(w, req) diff --git a/pkg/hold/pds/xrpc.go b/pkg/hold/pds/xrpc.go index 3f8da11..10897f5 100644 --- a/pkg/hold/pds/xrpc.go +++ b/pkg/hold/pds/xrpc.go @@ -479,6 +479,7 @@ func (h *XRPCHandler) HandleGetRecord(w http.ResponseWriter, r *http.Request) { // HandleListRecords lists records in a collection // Spec: https://docs.bsky.app/docs/api/com-atproto-repo-list-records // Supports pagination via limit, cursor, and reverse parameters +// Uses SQL index for efficient pagination (following official ATProto PDS pattern) func (h *XRPCHandler) HandleListRecords(w http.ResponseWriter, r *http.Request) { repoDID := r.URL.Query().Get("repo") collection := r.URL.Query().Get("collection") @@ -507,6 +508,95 @@ func (h *XRPCHandler) HandleListRecords(w http.ResponseWriter, r *http.Request) cursor := r.URL.Query().Get("cursor") reverse := r.URL.Query().Get("reverse") == "true" + // Use records index if available (efficient SQL-based pagination) + if h.pds.recordsIndex != nil { + h.handleListRecordsIndexed(w, r, collection, limit, cursor, reverse) + return + } + + // Fallback: MST-based listing (legacy path for tests or in-memory mode) + h.handleListRecordsMST(w, r, collection, limit, cursor, reverse) +} + +// handleListRecordsIndexed uses the SQL records index for efficient pagination +func (h *XRPCHandler) handleListRecordsIndexed(w http.ResponseWriter, r *http.Request, collection string, limit int, cursor string, reverse bool) { + // Query the index + indexedRecords, nextCursor, err := h.pds.recordsIndex.ListRecords(collection, limit, cursor, reverse) + if err != nil { + slog.Error("Failed to list records from index", "error", err, "collection", collection) + http.Error(w, fmt.Sprintf("failed to list records: %v", err), http.StatusInternalServerError) + return + } + + // Create session to fetch full record data + session, err := h.pds.carstore.ReadOnlySession(h.pds.uid) + if err != nil { + http.Error(w, fmt.Sprintf("failed to create session: %v", err), http.StatusInternalServerError) + return + } + + head, err := h.pds.carstore.GetUserRepoHead(r.Context(), h.pds.uid) + if err != nil { + http.Error(w, fmt.Sprintf("failed to get repo head: %v", err), http.StatusInternalServerError) + return + } + + if !head.Defined() { + // Empty repo, return empty list + response := map[string]any{"records": []any{}} + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(response) + return + } + + repoHandle, err := repo.OpenRepo(r.Context(), session, head) + if err != nil { + http.Error(w, fmt.Sprintf("failed to open repo: %v", err), http.StatusInternalServerError) + return + } + + // Fetch full record data for each indexed record + records := []map[string]any{} + for _, rec := range indexedRecords { + // Construct the record path + recordPath := rec.Collection + "/" + rec.Rkey + + // Get the record bytes + recordCID, recBytes, err := repoHandle.GetRecordBytes(r.Context(), recordPath) + if err != nil { + slog.Warn("Failed to get indexed record, skipping", "path", recordPath, "error", err) + continue + } + + // Decode using lexutil (type registry handles unmarshaling) + recordValue, err := lexutil.CborDecodeValue(*recBytes) + if err != nil { + slog.Warn("Failed to decode indexed record, skipping", "path", recordPath, "error", err) + continue + } + + records = append(records, map[string]any{ + "uri": fmt.Sprintf("at://%s/%s/%s", h.pds.DID(), rec.Collection, rec.Rkey), + "cid": recordCID.String(), + "value": recordValue, + }) + } + + response := map[string]any{ + "records": records, + } + + // Include cursor in response if there are more records + if nextCursor != "" { + response["cursor"] = nextCursor + } + + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(response) +} + +// handleListRecordsMST uses the legacy MST-based listing (fallback for tests) +func (h *XRPCHandler) handleListRecordsMST(w http.ResponseWriter, r *http.Request, collection string, limit int, cursor string, reverse bool) { // Generic implementation using repo.ForEach session, err := h.pds.carstore.ReadOnlySession(h.pds.uid) if err != nil { @@ -534,12 +624,11 @@ func (h *XRPCHandler) HandleListRecords(w http.ResponseWriter, r *http.Request) return } - // Initialize as empty slice (not nil) to ensure JSON encodes as [] not null - records := []map[string]any{} - var nextCursor string - skipUntilCursor := cursor != "" + // Collect all records in the collection first. + // MST only supports forward iteration, so for newest-first (default) we must + // collect all records, reverse, then apply cursor/limit. + allRecords := []map[string]any{} - // Iterate over all records in the collection err = repoHandle.ForEach(r.Context(), collection, func(k string, v cid.Cid) error { // k is like "io.atcr.hold.captain/self" or "io.atcr.hold.crew/3m3by7msdln22" parts := strings.Split(k, "/") @@ -552,27 +641,10 @@ func (h *XRPCHandler) HandleListRecords(w http.ResponseWriter, r *http.Request) rkey := parts[len(parts)-1] // Filter: only include records that match the requested collection - // MST keys are sorted lexicographically, so once we hit a different - // collection prefix, all remaining keys will also be outside our range if actualCollection != collection { return repo.ErrDoneIterating // Stop walking the tree } - // Handle cursor-based pagination - if skipUntilCursor { - if rkey == cursor { - skipUntilCursor = false // Found cursor, start including records after this - } - return nil // Skip this record - } - - // Check if we've hit the limit - if len(records) >= limit { - // Set next cursor to current rkey - nextCursor = rkey - return repo.ErrDoneIterating // Stop iteration - } - // Get the record bytes recordCID, recBytes, err := repoHandle.GetRecordBytes(r.Context(), k) if err != nil { @@ -585,41 +657,67 @@ func (h *XRPCHandler) HandleListRecords(w http.ResponseWriter, r *http.Request) return fmt.Errorf("failed to decode record: %v", err) } - records = append(records, map[string]any{ + allRecords = append(allRecords, map[string]any{ "uri": fmt.Sprintf("at://%s/%s/%s", h.pds.DID(), actualCollection, rkey), "cid": recordCID.String(), "value": recordValue, + "rkey": rkey, }) return nil }) if err != nil { - // ErrDoneIterating is expected when we stop walking early (reached collection boundary or hit limit) - // Check using strings.Contains because the error may be wrapped if err == repo.ErrDoneIterating || strings.Contains(err.Error(), "done iterating") { - // Successfully stopped at collection boundary or hit pagination limit, continue with collected records + // Successfully stopped at collection boundary } else if strings.Contains(err.Error(), "not found") { - // If the collection doesn't exist yet, return empty list - records = []map[string]any{} + allRecords = []map[string]any{} } else { http.Error(w, fmt.Sprintf("failed to list records: %v", err), http.StatusInternalServerError) return } } - // Default order is newest-first (reverse chronological), which requires - // reversing the MST's lexicographic order. When reverse=true, keep MST order. - if !reverse && len(records) > 0 { - for i, j := 0, len(records)-1; i < j; i, j = i+1, j-1 { - records[i], records[j] = records[j], records[i] + // Default order is newest-first (reverse chronological). + // MST iterates oldest-first, so reverse for default order. + if !reverse && len(allRecords) > 0 { + for i, j := 0, len(allRecords)-1; i < j; i, j = i+1, j-1 { + allRecords[i], allRecords[j] = allRecords[j], allRecords[i] } } + // Apply cursor and limit + records := []map[string]any{} + var nextCursor string + skipUntilCursor := cursor != "" + + for _, rec := range allRecords { + rkey := rec["rkey"].(string) + + if skipUntilCursor { + if rkey == cursor { + skipUntilCursor = false + } + continue + } + + if len(records) >= limit { + nextCursor = rkey + break + } + + delete(rec, "rkey") + records = append(records, rec) + } + + if skipUntilCursor { + records = []map[string]any{} + nextCursor = "" + } + response := map[string]any{ "records": records, } - // Include cursor in response if there are more records if nextCursor != "" { response["cursor"] = nextCursor } diff --git a/pkg/hold/pds/xrpc_test.go b/pkg/hold/pds/xrpc_test.go index a0d3cfd..225c4ec 100644 --- a/pkg/hold/pds/xrpc_test.go +++ b/pkg/hold/pds/xrpc_test.go @@ -29,6 +29,7 @@ import ( // setupTestXRPCHandler creates a fresh PDS instance and handler for each test // Bootstraps the PDS and suppresses logging to avoid log spam +// Uses :memory: database which disables RecordsIndex (uses MST fallback path) func setupTestXRPCHandler(t *testing.T) (*XRPCHandler, context.Context) { t.Helper() @@ -80,6 +81,73 @@ func setupTestXRPCHandler(t *testing.T) (*XRPCHandler, context.Context) { return handler, ctx } +// setupTestXRPCHandlerWithIndex creates a handler with file-based database +// to enable RecordsIndex (vs :memory: which disables it) +func setupTestXRPCHandlerWithIndex(t *testing.T) (*XRPCHandler, context.Context) { + t.Helper() + + ctx := context.Background() + tmpDir := t.TempDir() + + // Use file-based database to enable RecordsIndex + dbPath := filepath.Join(tmpDir, "pds.db") + keyPath := filepath.Join(tmpDir, "signing-key") + + // Copy shared signing key instead of generating a new one + if err := os.WriteFile(keyPath, sharedTestKey, 0600); err != nil { + t.Fatalf("Failed to copy shared signing key: %v", err) + } + + pds, err := NewHoldPDS(ctx, "did:web:hold.example.com", "https://hold.example.com", dbPath, keyPath, false) + if err != nil { + t.Fatalf("Failed to create test PDS: %v", err) + } + + // Verify RecordsIndex is enabled + if pds.RecordsIndex() == nil { + t.Fatal("Expected RecordsIndex to be non-nil for file-based database") + } + + // Bootstrap with a test owner, suppressing stdout to avoid log spam + ownerDID := "did:plc:testowner123" + + // Redirect stdout to suppress bootstrap logging + oldStdout := os.Stdout + r, w, _ := os.Pipe() + os.Stdout = w + + err = pds.Bootstrap(ctx, nil, ownerDID, true, false, "") + + // Restore stdout + w.Close() + os.Stdout = oldStdout + io.ReadAll(r) // Drain the pipe + + if err != nil { + t.Fatalf("Failed to bootstrap PDS: %v", err) + } + + // Wire up records indexing event handler + indexingHandler := pds.CreateRecordsIndexEventHandler(nil) + pds.RepomgrRef().SetEventHandler(indexingHandler, true) + + // Backfill index from MST (bootstrap created records but didn't index them) + if err := pds.BackfillRecordsIndex(ctx); err != nil { + t.Fatalf("Failed to backfill records index: %v", err) + } + + // Create mock PDS client for DPoP validation + mockClient := &mockPDSClient{} + + // Create mock s3 service and storage driver + mockS3 := s3.S3Service{} + + // Create XRPC handler with mock HTTP client + handler := NewXRPCHandler(pds, mockS3, nil, nil, mockClient) + + return handler, ctx +} + // Note: setupTestPDS is defined in captain_test.go and creates a PDS without bootstrapping // makeXRPCGetRequest creates a GET request with query parameters @@ -747,6 +815,234 @@ func TestHandleListRecords_MissingParameters(t *testing.T) { } } +// Tests for HandleListRecords with RecordsIndex (indexed path) +// These tests use file-based database to enable the SQL-based indexing + +// TestHandleListRecords_Indexed tests listing with RecordsIndex enabled +func TestHandleListRecords_Indexed(t *testing.T) { + handler, ctx := setupTestXRPCHandlerWithIndex(t) + holdDID := "did:web:hold.example.com" + + // Add crew members (will be indexed via event handler) + memberDIDs := []string{ + "did:plc:member1", + "did:plc:member2", + "did:plc:member3", + } + + for _, did := range memberDIDs { + _, err := handler.pds.AddCrewMember(ctx, did, "reader", []string{"blob:read"}) + if err != nil { + t.Fatalf("Failed to add crew member %s: %v", did, err) + } + } + + // Test listing crew records via indexed path + req := makeXRPCGetRequest(atproto.RepoListRecords, map[string]string{ + "repo": holdDID, + "collection": atproto.CrewCollection, + }) + w := httptest.NewRecorder() + + handler.HandleListRecords(w, req) + + result := assertJSONResponse(t, w, http.StatusOK) + + // Should have 4 crew records: 1 from bootstrap + 3 added + expectedCount := len(memberDIDs) + 1 + if records, ok := result["records"].([]any); !ok { + t.Error("Expected records array in response") + } else if len(records) != expectedCount { + t.Errorf("Expected %d crew records, got %d", expectedCount, len(records)) + } else { + // Verify each record has required fields + for i, rec := range records { + record, ok := rec.(map[string]any) + if !ok { + t.Errorf("Record %d: expected map, got %T", i, rec) + continue + } + + if uri, ok := record["uri"].(string); !ok || uri == "" { + t.Errorf("Record %d: expected uri string", i) + } + + if cid, ok := record["cid"].(string); !ok || cid == "" { + t.Errorf("Record %d: expected cid string", i) + } + + if value, ok := record["value"].(map[string]any); !ok { + t.Errorf("Record %d: expected value object", i) + } else { + if recordType, ok := value["$type"].(string); !ok || recordType != atproto.CrewCollection { + t.Errorf("Record %d: expected $type=%s, got %v", i, atproto.CrewCollection, value["$type"]) + } + } + } + } +} + +// TestHandleListRecords_Indexed_Pagination tests pagination with indexed path +func TestHandleListRecords_Indexed_Pagination(t *testing.T) { + handler, ctx := setupTestXRPCHandlerWithIndex(t) + holdDID := "did:web:hold.example.com" + + // Add 4 more crew members for total of 5 + for i := 0; i < 4; i++ { + _, err := handler.pds.AddCrewMember(ctx, fmt.Sprintf("did:plc:member%d", i), "reader", []string{"blob:read"}) + if err != nil { + t.Fatalf("Failed to add crew member: %v", err) + } + } + + // Test with limit=2 + req := makeXRPCGetRequest(atproto.RepoListRecords, map[string]string{ + "repo": holdDID, + "collection": atproto.CrewCollection, + "limit": "2", + }) + w := httptest.NewRecorder() + + handler.HandleListRecords(w, req) + + result := assertJSONResponse(t, w, http.StatusOK) + + // Verify we got exactly 2 records + records, ok := result["records"].([]any) + if !ok { + t.Fatal("Expected records array in response") + } + + if len(records) != 2 { + t.Errorf("Expected 2 records with limit=2, got %d", len(records)) + } + + // Verify cursor is present (there are more records) + cursor, ok := result["cursor"].(string) + if !ok || cursor == "" { + t.Fatal("Expected cursor in response when there are more records") + } + + // Test pagination with cursor + req2 := makeXRPCGetRequest(atproto.RepoListRecords, map[string]string{ + "repo": holdDID, + "collection": atproto.CrewCollection, + "limit": "2", + "cursor": cursor, + }) + w2 := httptest.NewRecorder() + + handler.HandleListRecords(w2, req2) + + result2 := assertJSONResponse(t, w2, http.StatusOK) + + records2, ok := result2["records"].([]any) + if !ok { + t.Fatal("Expected records array in paginated response") + } + + // Should get the next page of records + if len(records2) == 0 { + t.Error("Expected records in paginated response") + } + + // Verify no duplicates + seen := make(map[string]bool) + for _, r := range records { + rec := r.(map[string]any) + uri := rec["uri"].(string) + seen[uri] = true + } + for _, r := range records2 { + rec := r.(map[string]any) + uri := rec["uri"].(string) + if seen[uri] { + t.Errorf("Duplicate record in pagination: %s", uri) + } + } +} + +// TestHandleListRecords_Indexed_Reverse tests reverse ordering with indexed path +func TestHandleListRecords_Indexed_Reverse(t *testing.T) { + handler, ctx := setupTestXRPCHandlerWithIndex(t) + holdDID := "did:web:hold.example.com" + + // Add crew members + for i := 0; i < 3; i++ { + _, err := handler.pds.AddCrewMember(ctx, fmt.Sprintf("did:plc:member%d", i), "reader", []string{"blob:read"}) + if err != nil { + t.Fatalf("Failed to add crew member: %v", err) + } + } + + // Get normal order (default = newest first) + req1 := makeXRPCGetRequest(atproto.RepoListRecords, map[string]string{ + "repo": holdDID, + "collection": atproto.CrewCollection, + }) + w1 := httptest.NewRecorder() + handler.HandleListRecords(w1, req1) + result1 := assertJSONResponse(t, w1, http.StatusOK) + + // Get reverse order (oldest first) + req2 := makeXRPCGetRequest(atproto.RepoListRecords, map[string]string{ + "repo": holdDID, + "collection": atproto.CrewCollection, + "reverse": "true", + }) + w2 := httptest.NewRecorder() + handler.HandleListRecords(w2, req2) + result2 := assertJSONResponse(t, w2, http.StatusOK) + + records1 := result1["records"].([]any) + records2 := result2["records"].([]any) + + if len(records1) != len(records2) { + t.Fatalf("Expected same number of records, got %d vs %d", len(records1), len(records2)) + } + + if len(records1) > 1 { + // First record in normal order should be last in reverse order + first1 := records1[0].(map[string]any)["uri"].(string) + last2 := records2[len(records2)-1].(map[string]any)["uri"].(string) + + if first1 != last2 { + t.Error("Expected first record in default order to be last in reverse order") + } + } +} + +// TestHandleListRecords_Indexed_EmptyCollection tests empty collection with indexed path +func TestHandleListRecords_Indexed_EmptyCollection(t *testing.T) { + handler, _ := setupTestXRPCHandlerWithIndex(t) + holdDID := "did:web:hold.example.com" + + // List a collection that doesn't exist + req := makeXRPCGetRequest(atproto.RepoListRecords, map[string]string{ + "repo": holdDID, + "collection": "io.atcr.nonexistent", + }) + w := httptest.NewRecorder() + + handler.HandleListRecords(w, req) + + result := assertJSONResponse(t, w, http.StatusOK) + + records, ok := result["records"].([]any) + if !ok { + t.Fatal("Expected records array in response") + } + + if len(records) != 0 { + t.Errorf("Expected 0 records for empty collection, got %d", len(records)) + } + + // Should not have cursor for empty results + if _, ok := result["cursor"]; ok { + t.Error("Expected no cursor for empty collection") + } +} + // Tests for HandleDeleteRecord // TestHandleDeleteRecord tests com.atproto.repo.deleteRecord