diff --git a/internal/database/webhook_db_manager.go b/internal/database/webhook_db_manager.go index 7da8169..a628b56 100644 --- a/internal/database/webhook_db_manager.go +++ b/internal/database/webhook_db_manager.go @@ -41,6 +41,11 @@ type WebhookDBManager struct { dataDir string dbs sync.Map // map[webhookID]*gorm.DB log *slog.Logger + + // mu is held while a database is opened, deleted, or closed, so + // each file has at most one open handle. Reading an already cached + // handle does not take it. + mu sync.Mutex } // NewWebhookDBManager creates a new WebhookDBManager and @@ -86,43 +91,39 @@ func (m *WebhookDBManager) GetDB( ) (*gorm.DB, error) { // Fast path: already open if val, ok := m.dbs.Load(webhookID); ok { - cachedDB, castOK := val.(*gorm.DB) - if !castOK { - return nil, fmt.Errorf( - "%w for webhook %s", - errInvalidCachedDBType, - webhookID, - ) - } + return asGormDB(val, webhookID) + } - return cachedDB, nil + // Slow path: open the database under the lock, looking in the + // cache again first. A caller that raced another one here then + // waits for its handle instead of opening a second one. + m.mu.Lock() + defer m.mu.Unlock() + + if val, ok := m.dbs.Load(webhookID); ok { + return asGormDB(val, webhookID) } - // Slow path: open/create the database db, err := m.openDB(webhookID) if err != nil { return nil, err } - // Store it; if another goroutine beat us, close ours - actual, loaded := m.dbs.LoadOrStore(webhookID, db) - if loaded { - // Another goroutine created it first; close our duplicate - sqlDB, closeErr := db.DB() - if closeErr == nil { - _ = sqlDB.Close() - } + m.dbs.Store(webhookID, db) - existingDB, castOK := actual.(*gorm.DB) - if !castOK { - return nil, fmt.Errorf( - "%w for webhook %s", - errInvalidCachedDBType, - webhookID, - ) - } + return db, nil +} - return existingDB, nil +// asGormDB returns a value read from the cache as the database +// handle it is. +func asGormDB(val any, webhookID string) (*gorm.DB, error) { + db, ok := val.(*gorm.DB) + if !ok { + return nil, fmt.Errorf( + "%w for webhook %s", + errInvalidCachedDBType, + webhookID, + ) } return db, nil @@ -153,6 +154,11 @@ func (m *WebhookDBManager) DBExists( func (m *WebhookDBManager) DeleteDB( webhookID string, ) error { + // Held until the files are gone, so GetDB cannot open the file + // again between the close and the removal. + m.mu.Lock() + defer m.mu.Unlock() + // Close and remove from cache if val, ok := m.dbs.LoadAndDelete(webhookID); ok { if gormDB, castOK := val.(*gorm.DB); castOK { @@ -186,6 +192,11 @@ func (m *WebhookDBManager) DeleteDB( // CloseAll closes all open per-webhook database connections. // Called during application shutdown. func (m *WebhookDBManager) CloseAll() error { + // An open already under way finishes and is cached first, so it + // is closed here rather than cached after this loop has passed. + m.mu.Lock() + defer m.mu.Unlock() + var lastErr error m.dbs.Range(func(key, value any) bool { diff --git a/internal/database/webhook_db_manager_test.go b/internal/database/webhook_db_manager_test.go index 771f9e4..29a788c 100644 --- a/internal/database/webhook_db_manager_test.go +++ b/internal/database/webhook_db_manager_test.go @@ -1,10 +1,14 @@ package database_test import ( + "bytes" "context" + "log/slog" "net/http" "os" "path/filepath" + "strings" + "sync" "testing" "github.com/google/uuid" @@ -104,6 +108,54 @@ func TestWebhookDBManager_CreateAndGetDB(t *testing.T) { assert.Equal(t, `{"test": true}`, readEvent.Body) } +// Many callers ask for one webhook's database at the same moment, +// before it is cached. Only one of them may open the file; the others +// must wait for its handle. openDB logs one "opened per-webhook +// database" line per open, and those lines are what is counted. +func TestWebhookDBManager_ConcurrentFirstTouchOpensOnce(t *testing.T) { + t.Parallel() + + var logs bytes.Buffer + + mgr := database.NewTestWebhookDBManagerWithLogger( + t.TempDir(), + slog.New(slog.NewTextHandler(&logs, nil)), + ) + + t.Cleanup(func() { assert.NoError(t, mgr.CloseAll()) }) + + webhookID := uuid.New().String() + + const callers = 16 + + start := make(chan struct{}) + handles := make([]*gorm.DB, callers) + errs := make([]error, callers) + + var wg sync.WaitGroup + + for i := range callers { + wg.Go(func() { + <-start + + handles[i], errs[i] = mgr.GetDB(webhookID) + }) + } + + close(start) + wg.Wait() + + for i := range callers { + require.NoError(t, errs[i]) + assert.Same(t, handles[0], handles[i]) + } + + assert.Equal( + t, 1, + strings.Count(logs.String(), "opened per-webhook database"), + ) +} + func TestWebhookDBManager_DeleteDB(t *testing.T) { t.Parallel()