diff --git a/internal/config/config_validation_internal_test.go b/internal/config/config_validation_internal_test.go index 1800ab3..d3a3736 100644 --- a/internal/config/config_validation_internal_test.go +++ b/internal/config/config_validation_internal_test.go @@ -1,6 +1,7 @@ package config import ( + "database/sql" "log/slog" "os" "path/filepath" @@ -9,6 +10,8 @@ import ( "time" "git.eeqj.de/sneak/smartconfig" + + _ "modernc.org/sqlite" // SQLite driver registration ) // validTestSigningKey is a 32-character signing key that satisfies the @@ -94,12 +97,43 @@ func TestOmittedValuesUseDefaults(t *testing.T) { t.Errorf("AllowlistHosts = %v, want empty", c.AllowlistHosts) } - wantDBURL := "file:" + DefaultStateDir + "/state.sqlite3?_journal_mode=WAL" + wantDBURL := "file:" + DefaultStateDir + + "/state.sqlite3?_pragma=journal_mode(WAL)" if c.DBURL != wantDBURL { t.Errorf("DBURL = %q, want derived default %q", c.DBURL, wantDBURL) } } +// TestDefaultDBURLOpensTheDatabaseInWALMode opens the db_url derived from +// state_dir with the SQLite driver pixad uses and checks that the database +// is in WAL mode: the driver ignores any parameter it does not know. +func TestDefaultDBURLOpensTheDatabaseInWALMode(t *testing.T) { + t.Parallel() + + c, err := configFromYAML(t, signingKeyLine+"state_dir: "+t.TempDir()+"\n") + if err != nil { + t.Fatalf("config with only state_dir set should be valid, got: %v", err) + } + + db, err := sql.Open("sqlite", c.DBURL) + if err != nil { + t.Fatalf("failed to open %q: %v", c.DBURL, err) + } + + t.Cleanup(func() { _ = db.Close() }) + + var journalMode string + + err = db.QueryRowContext(t.Context(), "PRAGMA journal_mode").Scan(&journalMode) + if err != nil { + t.Fatalf("failed to read the journal mode of %q: %v", c.DBURL, err) + } + + if journalMode != "wal" { + t.Errorf("journal mode of %q = %q, want wal", c.DBURL, journalMode) + } +} + func TestExplicitValidValuesAreUsed(t *testing.T) { t.Parallel() diff --git a/internal/database/concurrent_writes_internal_test.go b/internal/database/concurrent_writes_internal_test.go new file mode 100644 index 0000000..3db4c4e --- /dev/null +++ b/internal/database/concurrent_writes_internal_test.go @@ -0,0 +1,153 @@ +package database + +import ( + "context" + "database/sql" + "fmt" + "log/slog" + "path/filepath" + "sync" + "testing" + + "sneak.berlin/go/pixa/internal/config" +) + +// TestConcurrentWritesAllSucceed opens a database the way pixad does and +// writes to it from several goroutines at once, so the writes run on +// separate connections, as one request's writes and the background eviction +// pass do. Every write must succeed, none failing with "database is locked", +// whether or not db_url already has parameters, and the parameters it has +// must still apply. +func TestConcurrentWritesAllSucceed(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + query string + wantJournalMode string + }{ + { + name: "db_url without parameters", + query: "", + wantJournalMode: "delete", + }, + { + name: "db_url with the WAL parameter", + query: "?_pragma=journal_mode(WAL)", + wantJournalMode: "wal", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + dbURL := "file:" + filepath.Join(t.TempDir(), "state.sqlite3") + tt.query + + d := &Database{ + log: slog.New(slog.DiscardHandler), + config: &config.Config{DBURL: dbURL}, + } + + err := d.connect(t.Context()) + if err != nil { + t.Fatalf("failed to connect to %q: %v", dbURL, err) + } + + t.Cleanup(func() { _ = d.db.Close() }) + + writeConcurrently(t, d.db) + + var journalMode string + + err = d.db.QueryRowContext(t.Context(), "PRAGMA journal_mode"). + Scan(&journalMode) + if err != nil { + t.Fatalf("failed to read the journal mode: %v", err) + } + + if journalMode != tt.wantJournalMode { + t.Errorf("journal mode = %q, want %q", journalMode, tt.wantJournalMode) + } + }) + } +} + +// writeConcurrently runs writeLikeOneRequest from several goroutines at once +// and checks that every write was made. +func writeConcurrently(t *testing.T, db *sql.DB) { + t.Helper() + + const ( + writers = 4 + requestsEach = 20 + totalRequests = writers * requestsEach + ) + + ctx := t.Context() + + var wg sync.WaitGroup + + for writer := range writers { + wg.Go(func() { + for request := range requestsEach { + key := fmt.Sprintf("%d-%d", writer, request) + + err := writeLikeOneRequest(ctx, db, key) + if err != nil { + t.Errorf("writer %d: %v", writer, err) + + return + } + } + }) + } + + wg.Wait() + + var hits, sources int + + err := db.QueryRowContext(ctx, ` + SELECT hit_count, (SELECT COUNT(*) FROM source_content) + FROM cache_stats WHERE id = 1 + `).Scan(&hits, &sources) + if err != nil { + t.Fatalf("failed to count the writes: %v", err) + } + + if hits != totalRequests || sources != totalRequests { + t.Errorf("hit_count = %d and %d source_content rows, want %d of each", + hits, sources, totalRequests) + } +} + +// writeLikeOneRequest makes the writes one request and the eviction pass +// make: it counts a cache hit, stores a source, records a transformed image +// and deletes that record again. +func writeLikeOneRequest(ctx context.Context, db *sql.DB, key string) error { + _, err := db.ExecContext(ctx, + `UPDATE cache_stats SET hit_count = hit_count + 1 WHERE id = 1`) + if err != nil { + return fmt.Errorf("counting a cache hit: %w", err) + } + + _, err = db.ExecContext(ctx, `INSERT INTO source_content + (content_hash, content_type, size_bytes) VALUES (?, 'image/png', 1)`, key) + if err != nil { + return fmt.Errorf("storing a source: %w", err) + } + + _, err = db.ExecContext(ctx, `INSERT INTO variant_content + (cache_key, size_bytes, content_type) VALUES (?, 1, 'image/png')`, key) + if err != nil { + return fmt.Errorf("recording a transformed image: %w", err) + } + + _, err = db.ExecContext(ctx, + `DELETE FROM variant_content WHERE cache_key = ?`, key) + if err != nil { + return fmt.Errorf("evicting a transformed image: %w", err) + } + + return nil +}