diff --git a/config.example.yaml b/config.example.yaml index 58a9de79..6316a37b 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -81,6 +81,10 @@ database: # Example: "postgres://user:password@localhost:5432/proxy?sslmode=disable" url: "" + # How often cache hit counts and last-access times are written, batched in + # one transaction. "0" writes each hit as it happens. + hit_flush_interval: "1s" + # Logging configuration log: # Minimum log level: "debug", "info", "warn", "error" diff --git a/docs/configuration.md b/docs/configuration.md index f09e0108..4e8983f7 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -97,6 +97,21 @@ database: |--------|-------------|------|-------------| | `database.url` | `PROXY_DATABASE_URL` | `-database-url` | PostgreSQL connection URL | +### Hit counts + +Every cache hit updates the artifact's hit count and last-access time, which the stats pages and LRU eviction use. Rather than writing each hit as its own transaction, the proxy counts hits in memory and writes them together every `hit_flush_interval`. SQLite allows one writer at a time, and the proxy uses a single SQLite connection, so with a write per hit, concurrent downloads queue behind each other. + +```yaml +database: + hit_flush_interval: "1s" +``` + +| Config | Environment | Flag | Description | +|--------|-------------|------|-------------| +| `database.hit_flush_interval` | `PROXY_DATABASE_HIT_FLUSH_INTERVAL` | - | How often batched hits are written (default `1s`). `0` writes each hit as it happens. | + +Hits not yet written are lost if the process is killed; a normal shutdown writes them. + ## Logging ```yaml diff --git a/internal/config/config.go b/internal/config/config.go index fcc41fa7..7ee5abfb 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -418,6 +418,12 @@ type DatabaseConfig struct { // URL is the PostgreSQL connection string. URL string `json:"url" yaml:"url"` + + // HitFlushInterval is how often cache hit counts and last-access times + // are written, batched in one transaction. Uses Go duration syntax + // (e.g. "1s", "10s"). Default: "1s". Set to "0" to write each hit as it + // happens. + HitFlushInterval string `json:"hit_flush_interval" yaml:"hit_flush_interval"` } // String returns a human-readable description of the configured database @@ -896,6 +902,7 @@ func (c *Config) LoadFromEnv() { setEnvString(&c.Database.Driver, "PROXY_DATABASE_DRIVER") setEnvString(&c.Database.Path, "PROXY_DATABASE_PATH") setEnvString(&c.Database.URL, "PROXY_DATABASE_URL") + setEnvString(&c.Database.HitFlushInterval, "PROXY_DATABASE_HIT_FLUSH_INTERVAL") setEnvString(&c.Log.Level, "PROXY_LOG_LEVEL") setEnvString(&c.Log.Format, "PROXY_LOG_FORMAT") setEnvString(&c.AccessLog.Path, "PROXY_ACCESS_LOG_PATH") @@ -1038,6 +1045,10 @@ func (c *Config) Validate() error { return err } + if err := validateHitFlushInterval(c.Database.HitFlushInterval); err != nil { + return err + } + return c.validateComponents() } @@ -1120,6 +1131,7 @@ const ( defaultMetadataTTL = 5 * time.Minute //nolint:mnd // sensible default defaultDirectServeTTL = 15 * time.Minute //nolint:mnd // sensible default defaultHTTPTimeout = 30 * time.Second //nolint:mnd // sensible default + defaultHitFlushInterval = time.Second defaultMetadataMaxSize = 100 << 20 defaultGradleBuildCacheMaxUploadSize = 100 << 20 defaultGradleBuildCacheSweepInterval = 10 * time.Minute @@ -1198,6 +1210,36 @@ func (c *Config) ParseHTTPTimeout() time.Duration { return d } +func validateHitFlushInterval(s string) error { + if s == "" || s == "0" { + return nil + } + d, err := time.ParseDuration(s) + if err != nil { + return fmt.Errorf("invalid database.hit_flush_interval %q: %w", s, err) + } + if d < 0 { + return fmt.Errorf("invalid database.hit_flush_interval %q: must be non-negative", s) + } + return nil +} + +// ParseHitFlushInterval returns how often batched cache hits are written. +// Returns 1 second if unset or invalid, 0 if explicitly disabled. +func (c *Config) ParseHitFlushInterval() time.Duration { + if c.Database.HitFlushInterval == "" { + return defaultHitFlushInterval + } + if c.Database.HitFlushInterval == "0" { + return 0 + } + d, err := time.ParseDuration(c.Database.HitFlushInterval) + if err != nil || d < 0 { + return defaultHitFlushInterval + } + return d +} + // ParseMetadataTTL returns the metadata TTL duration. // Returns 5 minutes if unset, 0 if explicitly disabled. func (c *Config) ParseMetadataTTL() time.Duration { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 050b3ce1..6a2bfcb7 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -939,6 +939,61 @@ func TestLoadHTTPTimeoutFromEnv(t *testing.T) { } } +func TestParseHitFlushInterval(t *testing.T) { + tests := []struct { + name string + interval string + want time.Duration + }{ + {"empty defaults to 1s", "", time.Second}, + {"explicit zero disables", "0", 0}, + {"10 seconds", "10s", 10 * time.Second}, + {"invalid defaults to 1s", "not-a-duration", time.Second}, + {"negative defaults to 1s", "-5s", time.Second}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := Default() + cfg.Database.HitFlushInterval = tt.interval + got := cfg.ParseHitFlushInterval() + if got != tt.want { + t.Errorf("ParseHitFlushInterval() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestValidateHitFlushInterval(t *testing.T) { + cfg := Default() + cfg.Database.HitFlushInterval = "not-a-duration" + if err := cfg.Validate(); err == nil { + t.Error("expected validation error for invalid hit_flush_interval") + } + + cfg.Database.HitFlushInterval = "-5s" + if err := cfg.Validate(); err == nil { + t.Error("expected validation error for negative hit_flush_interval") + } + + for _, ok := range []string{"10s", "0", ""} { + cfg.Database.HitFlushInterval = ok + if err := cfg.Validate(); err != nil { + t.Errorf("unexpected error for hit_flush_interval %q: %v", ok, err) + } + } +} + +func TestLoadHitFlushIntervalFromEnv(t *testing.T) { + cfg := Default() + t.Setenv("PROXY_DATABASE_HIT_FLUSH_INTERVAL", "10s") + cfg.LoadFromEnv() + + if cfg.Database.HitFlushInterval != "10s" { + t.Errorf("HitFlushInterval = %q, want %q", cfg.Database.HitFlushInterval, "10s") + } +} + func TestLoadMetadataTTLFromEnv(t *testing.T) { cfg := Default() t.Setenv("PROXY_METADATA_TTL", "10m") diff --git a/internal/database/database.go b/internal/database/database.go index 6b326eca..f438f676 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -36,6 +36,7 @@ type DB struct { *sqlx.DB dialect Dialect path string + hits *hitBatch } func (db *DB) Dialect() Dialect { diff --git a/internal/database/hits.go b/internal/database/hits.go new file mode 100644 index 00000000..249c6633 --- /dev/null +++ b/internal/database/hits.go @@ -0,0 +1,146 @@ +package database + +import ( + "cmp" + "log/slog" + "slices" + "sync" + "time" +) + +// Recording each cache hit as its own UPDATE makes every cached download a +// write transaction, and SQLite runs on a single connection, so concurrent +// downloads queue behind those writes. BatchHits counts hits in memory and +// writes them in one transaction per interval instead. + +type hitKey struct{ versionPURL, filename string } + +type hitEntry struct { + count int64 + last time.Time +} + +type hitBatch struct { + mu sync.Mutex + pending map[hitKey]hitEntry + stop chan struct{} + stopOnce sync.Once + done chan struct{} +} + +func (b *hitBatch) add(k hitKey, e hitEntry) { + b.mu.Lock() + defer b.mu.Unlock() + cur := b.pending[k] + cur.count += e.count + if e.last.After(cur.last) { + cur.last = e.last + } + b.pending[k] = cur +} + +func (b *hitBatch) take() map[hitKey]hitEntry { + b.mu.Lock() + defer b.mu.Unlock() + pending := b.pending + b.pending = map[hitKey]hitEntry{} + return pending +} + +// BatchHits makes RecordArtifactHit buffer hits and write them every +// interval. An interval of zero or less leaves every hit written immediately. +// Close writes whatever is still pending. Failed writes are logged to logger. +func (db *DB) BatchHits(interval time.Duration, logger *slog.Logger) { + if interval <= 0 || db.hits != nil { + return + } + if logger == nil { + logger = slog.New(slog.DiscardHandler) + } + b := &hitBatch{pending: map[hitKey]hitEntry{}, stop: make(chan struct{}), done: make(chan struct{})} + db.hits = b + go func() { + defer close(b.done) + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ticker.C: + if err := db.flushHits(); err != nil { + logger.Warn("failed to write cache hits, retrying next flush", "error", err) + } + case <-b.stop: + if err := db.flushHits(); err != nil { + logger.Error("failed to write cache hits on close, dropping them", "error", err) + } + return + } + } + }() +} + +// flushHits writes the pending hits in one transaction. If the transaction +// fails they are put back for the next flush, so a failed write delays +// counts rather than losing them. +func (db *DB) flushHits() error { + pending := db.hits.take() + if len(pending) == 0 { + return nil + } + + err := db.writeHits(pending) + if err != nil { + for k, e := range pending { + db.hits.add(k, e) + } + } + return err +} + +func (db *DB) writeHits(pending map[hitKey]hitEntry) error { + tx, err := db.Beginx() + if err != nil { + return err + } + defer func() { _ = tx.Rollback() }() + + // Timestamps only move forward: with proxies sharing a database, an older + // batch can flush after a newer hit is already written. + stmt, err := tx.Preparex(db.Rebind(` + UPDATE artifacts + SET hit_count = hit_count + ?, + last_accessed_at = CASE WHEN last_accessed_at IS NULL OR last_accessed_at < ? THEN ? ELSE last_accessed_at END, + updated_at = CASE WHEN updated_at IS NULL OR updated_at < ? THEN ? ELSE updated_at END + WHERE version_purl = ? AND filename = ? + `)) + if err != nil { + return err + } + defer func() { _ = stmt.Close() }() + + // A fixed order keeps proxies sharing a Postgres database from locking + // the same rows in opposite orders and deadlocking. + keys := make([]hitKey, 0, len(pending)) + for k := range pending { + keys = append(keys, k) + } + slices.SortFunc(keys, func(a, b hitKey) int { + return cmp.Or(cmp.Compare(a.versionPURL, b.versionPURL), cmp.Compare(a.filename, b.filename)) + }) + for _, k := range keys { + e := pending[k] + if _, err := stmt.Exec(e.count, e.last, e.last, e.last, e.last, k.versionPURL, k.filename); err != nil { + return err + } + } + return tx.Commit() +} + +// Close writes any batched hits, then closes the database. +func (db *DB) Close() error { + if b := db.hits; b != nil { + b.stopOnce.Do(func() { close(b.stop) }) + <-b.done + } + return db.DB.Close() +} diff --git a/internal/database/hits_test.go b/internal/database/hits_test.go new file mode 100644 index 00000000..99126f6b --- /dev/null +++ b/internal/database/hits_test.go @@ -0,0 +1,246 @@ +package database + +import ( + "bytes" + "log/slog" + "path/filepath" + "strings" + "sync" + "testing" + "time" +) + +var discardLogger = slog.New(slog.DiscardHandler) + +func seedHitTestArtifact(t *testing.T, db *DB) (string, string) { + t.Helper() + versionPURL, filename := "pkg:npm/lodash@4.17.21", "lodash-4.17.21.tgz" + if err := db.UpsertPackage(&Package{PURL: "pkg:npm/lodash", Ecosystem: "npm", Name: "lodash"}); err != nil { + t.Fatalf("UpsertPackage failed: %v", err) + } + if err := db.UpsertVersion(&Version{PURL: versionPURL, PackagePURL: "pkg:npm/lodash"}); err != nil { + t.Fatalf("UpsertVersion failed: %v", err) + } + if err := db.UpsertArtifact(&Artifact{ + VersionPURL: versionPURL, + Filename: filename, + UpstreamURL: "https://registry.npmjs.org/lodash/-/" + filename, + }); err != nil { + t.Fatalf("UpsertArtifact failed: %v", err) + } + return versionPURL, filename +} + +func hitCount(t *testing.T, db *DB, versionPURL, filename string) int64 { + t.Helper() + a, err := db.GetArtifact(versionPURL, filename) + if err != nil || a == nil { + t.Fatalf("GetArtifact failed: %v", err) + } + return a.HitCount +} + +func TestBatchHitsWritesOnFlush(t *testing.T) { + runWithBothDatabases(t, func(t *testing.T, db *DB) { + versionPURL, filename := seedHitTestArtifact(t, db) + db.BatchHits(time.Hour, discardLogger) + + for range 3 { + if err := db.RecordArtifactHit(versionPURL, filename); err != nil { + t.Fatalf("RecordArtifactHit failed: %v", err) + } + } + if got := hitCount(t, db, versionPURL, filename); got != 0 { + t.Fatalf("hit count before flush = %d, want 0", got) + } + + if err := db.flushHits(); err != nil { + t.Fatalf("flushHits failed: %v", err) + } + a, err := db.GetArtifact(versionPURL, filename) + if err != nil { + t.Fatalf("GetArtifact failed: %v", err) + } + if a.HitCount != 3 { + t.Errorf("hit count after flush = %d, want 3", a.HitCount) + } + if !a.LastAccessedAt.Valid { + t.Error("expected last_accessed_at to be set") + } + }) +} + +func TestBatchHitsWritesOnClose(t *testing.T) { + path := filepath.Join(t.TempDir(), "test.db") + db, err := Create(path) + if err != nil { + t.Fatalf("Create failed: %v", err) + } + versionPURL, filename := seedHitTestArtifact(t, db) + db.BatchHits(time.Hour, discardLogger) + + for range 2 { + _ = db.RecordArtifactHit(versionPURL, filename) + } + if err := db.Close(); err != nil { + t.Fatalf("Close failed: %v", err) + } + + db, err = Open(path) + if err != nil { + t.Fatalf("Open failed: %v", err) + } + defer func() { _ = db.Close() }() + if got := hitCount(t, db, versionPURL, filename); got != 2 { + t.Errorf("hit count after Close = %d, want 2", got) + } +} + +// TestBatchHitsKeepNewerAccessTime flushes a batch older than a hit already +// written, as when proxies share a database. The count still adds up, but +// neither timestamp may move backwards. +func TestBatchHitsKeepNewerAccessTime(t *testing.T) { + runWithBothDatabases(t, func(t *testing.T, db *DB) { + versionPURL, filename := seedHitTestArtifact(t, db) + k := hitKey{versionPURL, filename} + now := time.Now() + read := func() *Artifact { + t.Helper() + a, err := db.GetArtifact(versionPURL, filename) + if err != nil || a == nil { + t.Fatalf("GetArtifact failed: %v", err) + } + return a + } + write := func(last time.Time) { + t.Helper() + if err := db.writeHits(map[hitKey]hitEntry{k: {count: 1, last: last}}); err != nil { + t.Fatalf("writeHits failed: %v", err) + } + } + + write(now) + newer := read() + if !newer.LastAccessedAt.Valid { + t.Fatal("last_accessed_at not set from NULL") + } + + write(now.Add(-time.Minute)) + got := read() + if got.HitCount != 2 { + t.Errorf("hit count = %d, want 2", got.HitCount) + } + if !got.LastAccessedAt.Time.Equal(newer.LastAccessedAt.Time) { + t.Errorf("last_accessed_at moved from %v to %v", newer.LastAccessedAt.Time, got.LastAccessedAt.Time) + } + if !got.UpdatedAt.Equal(newer.UpdatedAt) { + t.Errorf("updated_at moved from %v to %v", newer.UpdatedAt, got.UpdatedAt) + } + + write(now.Add(time.Minute)) + got = read() + if !got.LastAccessedAt.Time.After(newer.LastAccessedAt.Time) || !got.UpdatedAt.After(newer.UpdatedAt) { + t.Error("a newer batch did not advance the timestamps") + } + }) +} + +func TestBatchHitsZeroIntervalWritesImmediately(t *testing.T) { + runWithBothDatabases(t, func(t *testing.T, db *DB) { + versionPURL, filename := seedHitTestArtifact(t, db) + db.BatchHits(0, discardLogger) + + _ = db.RecordArtifactHit(versionPURL, filename) + if got := hitCount(t, db, versionPURL, filename); got != 1 { + t.Errorf("hit count = %d, want 1", got) + } + }) +} + +func TestBatchHitsKeepsHitsWhenWriteFails(t *testing.T) { + db := createTestDB(t) + versionPURL, filename := seedHitTestArtifact(t, db) + db.BatchHits(time.Hour, discardLogger) + + for range 2 { + _ = db.RecordArtifactHit(versionPURL, filename) + } + _ = db.DB.Close() + if err := db.flushHits(); err == nil { + t.Fatal("expected flushHits to fail on a closed database") + } + if got := db.hits.take()[hitKey{versionPURL, filename}].count; got != 2 { + t.Errorf("pending hits after failed flush = %d, want 2", got) + } + _ = db.Close() +} + +// TestBatchHitsCountsEveryHitWhileFlushing records hits from several goroutines +// while another flushes, so no hit may be lost between taking the pending +// hits and writing them. +func TestBatchHitsCountsEveryHitWhileFlushing(t *testing.T) { + runWithBothDatabases(t, func(t *testing.T, db *DB) { + versionPURL, filename := seedHitTestArtifact(t, db) + db.BatchHits(time.Hour, discardLogger) + + const writers, hitsEach = 8, 100 + stop := make(chan struct{}) + flushed := make(chan struct{}) + go func() { + defer close(flushed) + for { + select { + case <-stop: + return + default: + if err := db.flushHits(); err != nil { + t.Errorf("flushHits failed: %v", err) + return + } + } + } + }() + var wg sync.WaitGroup + for range writers { + wg.Go(func() { + for range hitsEach { + if err := db.RecordArtifactHit(versionPURL, filename); err != nil { + t.Errorf("RecordArtifactHit failed: %v", err) + } + } + }) + } + wg.Wait() + close(stop) + <-flushed + + if err := db.flushHits(); err != nil { + t.Fatalf("final flushHits failed: %v", err) + } + if got := hitCount(t, db, versionPURL, filename); got != writers*hitsEach { + t.Errorf("hit count = %d, want %d", got, writers*hitsEach) + } + }) +} + +func TestBatchHitsLogsHitsDroppedOnClose(t *testing.T) { + db := createTestDB(t) + versionPURL, filename := seedHitTestArtifact(t, db) + var logs bytes.Buffer + db.BatchHits(time.Hour, slog.New(slog.NewTextHandler(&logs, nil))) + + _ = db.RecordArtifactHit(versionPURL, filename) + _ = db.DB.Close() + _ = db.Close() + + if !strings.Contains(logs.String(), "dropping them") { + t.Errorf("no error logged for hits dropped on close, logs: %q", logs.String()) + } +} + +func TestCloseTwiceWithBatchHits(t *testing.T) { + db := createTestDB(t) + db.BatchHits(time.Hour, discardLogger) + _ = db.Close() + _ = db.Close() +} diff --git a/internal/database/queries.go b/internal/database/queries.go index 12d076e2..4e568d69 100644 --- a/internal/database/queries.go +++ b/internal/database/queries.go @@ -398,6 +398,10 @@ func (db *DB) upsertArtifactFrom(a *Artifact, previous sql.NullString) (bool, er func (db *DB) RecordArtifactHit(versionPURL, filename string) error { now := time.Now() + if db.hits != nil { + db.hits.add(hitKey{versionPURL, filename}, hitEntry{count: 1, last: now}) + return nil + } query := db.Rebind(` UPDATE artifacts SET hit_count = hit_count + 1, last_accessed_at = ?, updated_at = ? diff --git a/internal/server/server.go b/internal/server/server.go index 35066486..648a3b8f 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -151,6 +151,7 @@ func New(cfg *config.Config, logger *slog.Logger, buildInfo BuildInfo) (*Server, _ = db.Close() return nil, fmt.Errorf("migrating database schema: %w", err) } + db.BatchHits(cfg.ParseHitFlushInterval(), logger) // Initialize storage storageURL := cfg.Storage.URL