From 375b472143f1c227dd6e0eec865b4ad0785c16c2 Mon Sep 17 00:00:00 2001 From: Syngnat Date: Wed, 22 Jul 2026 07:47:34 +0800 Subject: [PATCH] =?UTF-8?q?=E2=9A=A1=EF=B8=8F=20perf(db-cache):=20?= =?UTF-8?q?=E9=98=B2=E6=AD=A2=E9=87=8A=E6=94=BE=E7=AB=9E=E6=80=81=E5=A4=8D?= =?UTF-8?q?=E6=B4=BB=E6=95=B0=E6=8D=AE=E5=BA=93=E8=BF=9E=E6=8E=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/app/app.go | 352 ++++++++-- internal/app/app_db_cache_concurrency_test.go | 657 ++++++++++++++++++ 2 files changed, 969 insertions(+), 40 deletions(-) diff --git a/internal/app/app.go b/internal/app/app.go index e353ce4c..2d26c893 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -50,6 +50,11 @@ var ( defaultAppTextLocalizer *i18n.Localizer ) +var ( + errDatabaseConnectionReleased = errors.New("数据库连接请求已被释放") + errDatabaseConnectionShutdown = errors.New("应用正在关闭,无法建立数据库连接") +) + type cachedDatabase struct { inst db.Database lastPing time.Time @@ -114,6 +119,19 @@ type cachedConnectFailure struct { err error } +type databaseConnectFlight struct { + id uint64 + groupKey string + cacheKey string + releaseMatchKey string + cancelErr error +} + +type databaseConnectResult struct { + inst db.Database + cacheKey string +} + type queryContext struct { cancel context.CancelFunc started time.Time @@ -142,6 +160,10 @@ type App struct { dbCache map[string]cachedDatabase // Cache for DB connections connectFailures map[string]cachedConnectFailure dbConnectGroup singleflight.Group + dbConnectFlights map[uint64]*databaseConnectFlight + nextDBConnectFlightID uint64 + dbShuttingDown bool + dbConnectBeforeForgetHook func() // Test seam for release/singleflight ordering. mu sync.RWMutex // Mutex for cache access updateMu sync.Mutex updateState updateState @@ -201,6 +223,7 @@ func NewAppWithSecretStore(store secretstore.SecretStore) *App { return &App{ dbCache: make(map[string]cachedDatabase), connectFailures: make(map[string]cachedConnectFailure), + dbConnectFlights: make(map[uint64]*databaseConnectFlight), runningQueries: make(map[string]queryContext), sqlTransactions: make(map[string]*managedSQLTransaction), configDir: resolveAppConfigDir(), @@ -380,17 +403,11 @@ func (a *App) LogWindowDiagnostic(stage string, payload string) { // Shutdown is called when the app terminates. func (a *App) Shutdown() { logger.Infof("应用开始关闭,准备释放资源") + a.beginDatabaseShutdown() a.stopConnectionKeepAliveLoop() a.rollbackPendingSQLTransactionsOnShutdown() a.closeSQLAuditStore() - a.mu.Lock() - defer a.mu.Unlock() - for _, dbInst := range a.dbCache { - if err := dbInst.inst.Close(); err != nil { - logger.Error(err, "关闭数据库连接失败") - } - } - a.dbCache = make(map[string]cachedDatabase) + a.closeCachedDatabasesForShutdown() proxytunnel.CloseAllForwarders() // Close all Redis connections CloseAllRedisClients() @@ -517,20 +534,185 @@ type cachedDatabaseCloseTarget struct { inst db.Database } +func (a *App) beginDatabaseConnectFlight(groupKey string, config connection.ConnectionConfig) (*databaseConnectFlight, error) { + if a == nil { + return nil, errDatabaseConnectionShutdown + } + + a.mu.Lock() + defer a.mu.Unlock() + if a.dbShuttingDown { + return nil, errDatabaseConnectionShutdown + } + if a.dbConnectFlights == nil { + a.dbConnectFlights = make(map[uint64]*databaseConnectFlight) + } + // Keep only active physical leaders. Release can invalidate these tokens + // without retaining a generation/tombstone for every connection ever seen. + a.nextDBConnectFlightID++ + flight := &databaseConnectFlight{ + id: a.nextDBConnectFlightID, + groupKey: groupKey, + cacheKey: groupKey, + releaseMatchKey: getConnectionReleaseMatchKey(config), + } + a.dbConnectFlights[flight.id] = flight + return flight, nil +} + +func (a *App) finishDatabaseConnectFlight(flight *databaseConnectFlight) { + if a == nil || flight == nil { + return + } + a.mu.Lock() + if current, exists := a.dbConnectFlights[flight.id]; exists && current == flight { + delete(a.dbConnectFlights, flight.id) + } + a.mu.Unlock() +} + +func (a *App) databaseConnectFlightErrorLocked(flight *databaseConnectFlight) error { + if a.dbShuttingDown { + return errDatabaseConnectionShutdown + } + if flight == nil { + return errDatabaseConnectionReleased + } + current, exists := a.dbConnectFlights[flight.id] + if !exists || current != flight { + return errDatabaseConnectionReleased + } + return flight.cancelErr +} + +func (a *App) databaseConnectFlightError(flight *databaseConnectFlight) error { + if a == nil { + return errDatabaseConnectionShutdown + } + a.mu.RLock() + defer a.mu.RUnlock() + return a.databaseConnectFlightErrorLocked(flight) +} + +func (a *App) databaseConnectionReturnError(cacheKey string, inst db.Database) error { + if a == nil { + return errDatabaseConnectionShutdown + } + a.mu.RLock() + defer a.mu.RUnlock() + if a.dbShuttingDown { + return errDatabaseConnectionShutdown + } + entry, exists := a.dbCache[cacheKey] + if !exists || entry.inst == nil || entry.inst != inst { + return errDatabaseConnectionReleased + } + return nil +} + +func (a *App) cancelDatabaseConnectFlightsLocked(match func(*databaseConnectFlight) bool, cancelErr error, excludedFlightID uint64) []string { + groupKeys := make([]string, 0) + for _, flight := range a.dbConnectFlights { + if flight == nil || flight.id == excludedFlightID || !match(flight) { + continue + } + if flight.cancelErr == nil || errors.Is(cancelErr, errDatabaseConnectionShutdown) { + flight.cancelErr = cancelErr + } + if a.connectFailures != nil { + delete(a.connectFailures, flight.cacheKey) + } + groupKeys = append(groupKeys, flight.groupKey) + } + return groupKeys +} + +func (a *App) forgetDatabaseConnectGroupsLocked(groupKeys []string) { + if a == nil || len(groupKeys) == 0 { + return + } + // Keep Forget in the same app-cache critical section as flight cancellation. + // Otherwise an old leader can finish, a fresh group can be installed, and a + // delayed Forget can accidentally remove that fresh group (ABA). + if a.dbConnectBeforeForgetHook != nil { + a.dbConnectBeforeForgetHook() + } + forgotten := make(map[string]struct{}, len(groupKeys)) + for _, groupKey := range groupKeys { + if _, exists := forgotten[groupKey]; exists { + continue + } + forgotten[groupKey] = struct{}{} + // A request that starts after release must create a fresh physical flight + // instead of joining the invalidated leader still unwinding in Connect. + a.dbConnectGroup.Forget(groupKey) + } +} + +func (a *App) beginDatabaseShutdown() { + if a == nil { + return + } + a.mu.Lock() + a.dbShuttingDown = true + groupKeys := a.cancelDatabaseConnectFlightsLocked(func(*databaseConnectFlight) bool { return true }, errDatabaseConnectionShutdown, 0) + a.forgetDatabaseConnectGroupsLocked(groupKeys) + a.mu.Unlock() +} + +func (a *App) closeCachedDatabasesForShutdown() { + if a == nil { + return + } + targets := make([]cachedDatabaseCloseTarget, 0) + a.mu.Lock() + for key, entry := range a.dbCache { + targets = append(targets, cachedDatabaseCloseTarget{key: key, inst: entry.inst}) + } + a.dbCache = make(map[string]cachedDatabase) + a.connectFailures = make(map[string]cachedConnectFailure) + a.mu.Unlock() + + for _, target := range targets { + if target.inst == nil { + continue + } + if err := target.inst.Close(); err != nil { + logger.Error(err, "关闭数据库连接失败:缓存Key=%s", shortCacheKey(target.key)) + } + } +} + func (a *App) releaseCachedDatabaseConnectionsForConfig(config connection.ConnectionConfig) int { if a == nil { return 0 } - return a.releaseCachedDatabaseConnectionsByMatchKey(getConnectionReleaseMatchKey(config)) + return a.releaseCachedDatabaseConnectionsForConfigExcludingFlight(config, 0) +} + +func (a *App) releaseCachedDatabaseConnectionsForConfigExcludingFlight(config connection.ConnectionConfig, excludedFlightID uint64) int { + if a == nil { + return 0 + } + return a.releaseCachedDatabaseConnectionsByMatchKeyExcludingFlight(getConnectionReleaseMatchKey(config), excludedFlightID) } func (a *App) releaseCachedDatabaseConnectionsByMatchKey(targetKey string) int { + return a.releaseCachedDatabaseConnectionsByMatchKeyExcludingFlight(targetKey, 0) +} + +func (a *App) releaseCachedDatabaseConnectionsByMatchKeyExcludingFlight(targetKey string, excludedFlightID uint64) int { if a == nil || strings.TrimSpace(targetKey) == "" { return 0 } targets := make([]cachedDatabaseCloseTarget, 0) a.mu.Lock() + // Mark leaders under the same lock used by the final cache write. This is + // the release/store linearization point that prevents late resurrection. + groupKeys := a.cancelDatabaseConnectFlightsLocked(func(flight *databaseConnectFlight) bool { + return flight.releaseMatchKey == targetKey + }, errDatabaseConnectionReleased, excludedFlightID) for key, entry := range a.dbCache { entryConfig := entry.config if strings.TrimSpace(entryConfig.Type) == "" { @@ -540,8 +722,10 @@ func (a *App) releaseCachedDatabaseConnectionsByMatchKey(targetKey string) int { continue } targets = append(targets, cachedDatabaseCloseTarget{key: key, inst: entry.inst}) + groupKeys = append(groupKeys, key) delete(a.dbCache, key) } + a.forgetDatabaseConnectGroupsLocked(groupKeys) a.mu.Unlock() for _, target := range targets { @@ -623,12 +807,18 @@ func (a *App) invalidateCachedDatabase(config connection.ConnectionConfig, reaso shortKey := shortCacheKey(key) a.mu.Lock() + groupKeys := a.cancelDatabaseConnectFlightsLocked(func(flight *databaseConnectFlight) bool { + return flight.cacheKey == key + }, errDatabaseConnectionReleased, 0) + groupKeys = append(groupKeys, key) entry, exists := a.dbCache[key] if !exists || entry.inst == nil { + a.forgetDatabaseConnectGroupsLocked(groupKeys) a.mu.Unlock() return false } delete(a.dbCache, key) + a.forgetDatabaseConnectGroupsLocked(groupKeys) a.mu.Unlock() if closeErr := entry.inst.Close(); closeErr != nil { @@ -942,6 +1132,12 @@ func (a *App) getDatabaseWithPing(config connection.ConnectionConfig, forcePing if err != nil { return nil, err } + a.mu.RLock() + shuttingDown := a.dbShuttingDown + a.mu.RUnlock() + if shuttingDown { + return nil, errDatabaseConnectionShutdown + } isFileDB := isFileDatabaseType(effectiveConfig.Type) key := getCacheKey(effectiveConfig) @@ -960,10 +1156,15 @@ func (a *App) getDatabaseWithPing(config connection.ConnectionConfig, forcePing // Best-effort cleanup: if cached instance exists for this exact config, close it. var staleDatabase db.Database a.mu.Lock() + groupKeys := a.cancelDatabaseConnectFlightsLocked(func(flight *databaseConnectFlight) bool { + return flight.cacheKey == key + }, errDatabaseConnectionReleased, 0) + groupKeys = append(groupKeys, key) if cur, exists := a.dbCache[key]; exists && cur.inst != nil { staleDatabase = cur.inst delete(a.dbCache, key) } + a.forgetDatabaseConnectGroupsLocked(groupKeys) a.mu.Unlock() if staleDatabase != nil { _ = staleDatabase.Close() @@ -996,6 +1197,9 @@ func (a *App) getDatabaseWithPing(config connection.ConnectionConfig, forcePing if isFileDB { logger.Infof("复用文件库连接缓存(免 Ping):类型=%s 缓存Key=%s", strings.TrimSpace(effectiveConfig.Type), shortKey) } + if returnErr := a.databaseConnectionReturnError(key, entry.inst); returnErr != nil { + return nil, returnErr + } return entry.inst, nil } @@ -1010,6 +1214,9 @@ func (a *App) getDatabaseWithPing(config connection.ConnectionConfig, forcePing if isFileDB { logger.Infof("复用文件库连接缓存(Ping 成功):类型=%s 缓存Key=%s", strings.TrimSpace(effectiveConfig.Type), shortKey) } + if returnErr := a.databaseConnectionReturnError(key, entry.inst); returnErr != nil { + return nil, returnErr + } return entry.inst, nil } else { logger.Error(err, "缓存连接不可用,准备重建:%s 缓存Key=%s", formatConnSummary(effectiveConfig), shortKey) @@ -1039,17 +1246,30 @@ func (a *App) getDatabaseWithPing(config connection.ConnectionConfig, forcePing return nil, failureErr } value, err, _ := a.dbConnectGroup.Do(key, func() (any, error) { - return a.connectAndCacheDatabase(effectiveConfig, key, isFileDB) + flight, beginErr := a.beginDatabaseConnectFlight(key, effectiveConfig) + if beginErr != nil { + return nil, beginErr + } + defer a.finishDatabaseConnectFlight(flight) + return a.connectAndCacheDatabase(effectiveConfig, key, isFileDB, flight) }) if err != nil { return nil, err } - dbInst, ok := value.(db.Database) - if !ok || dbInst == nil { + result, ok := value.(databaseConnectResult) + if !ok || result.inst == nil || strings.TrimSpace(result.cacheKey) == "" { return nil, fmt.Errorf("数据库连接缓存返回了无效实例") } - a.applyCachedDatabaseKeepAlivePolicy(key, dbInst, resolveConnectionKeepAlivePolicy(effectiveConfig), time.Now()) - return dbInst, nil + if _, exists := a.applyCachedDatabaseKeepAlivePolicy(result.cacheKey, result.inst, resolveConnectionKeepAlivePolicy(effectiveConfig), time.Now()); !exists { + if returnErr := a.databaseConnectionReturnError(result.cacheKey, result.inst); returnErr != nil { + return nil, returnErr + } + return nil, errDatabaseConnectionReleased + } + if returnErr := a.databaseConnectionReturnError(result.cacheKey, result.inst); returnErr != nil { + return nil, returnErr + } + return result.inst, nil } func (a *App) applyCachedDatabaseKeepAlivePolicy(key string, expectedInst db.Database, policy connectionKeepAlivePolicy, now time.Time) (cachedDatabase, bool) { @@ -1067,7 +1287,7 @@ func (a *App) applyCachedDatabaseKeepAlivePolicy(key string, expectedInst db.Dat return entry, true } -func (a *App) connectAndCacheDatabase(effectiveConfig connection.ConnectionConfig, initialKey string, isFileDB bool) (db.Database, error) { +func (a *App) connectAndCacheDatabase(effectiveConfig connection.ConnectionConfig, initialKey string, isFileDB bool, flight *databaseConnectFlight) (databaseConnectResult, error) { key := initialKey shortKey := shortenCacheKey(key) @@ -1075,50 +1295,77 @@ func (a *App) connectAndCacheDatabase(effectiveConfig connection.ConnectionConfi // finishes its keyed connection flight. Recheck after becoming the leader so // a completed flight can never be followed by a duplicate physical connect. a.mu.RLock() + flightErr := a.databaseConnectFlightErrorLocked(flight) existing, exists := a.dbCache[key] a.mu.RUnlock() + if flightErr != nil { + return databaseConnectResult{}, flightErr + } if exists && existing.inst != nil { - return existing.inst, nil + return databaseConnectResult{inst: existing.inst, cacheKey: key}, nil } if failureErr := a.cachedConnectFailureError(effectiveConfig, key, "db.backend.message.connect_failure_cooldown"); failureErr != nil { - return nil, failureErr + return databaseConnectResult{}, failureErr } if revisionErr := verifyDriverAgentRevisionFunc(effectiveConfig); revisionErr != nil { - return nil, withLogHint{err: revisionErr, logPath: logger.Path()} + return databaseConnectResult{}, withLogHint{err: revisionErr, logPath: logger.Path()} } dbInst, connectedConfig, err := a.connectEffectiveDatabaseWithStartupRetry(effectiveConfig) + if flightErr := a.databaseConnectFlightError(flight); flightErr != nil { + if dbInst != nil { + _ = dbInst.Close() + } + return databaseConnectResult{}, flightErr + } if err != nil { - retryInst, retryConfig, retryErr := a.retryConnectAfterMySQLMaxUserConnections(effectiveConfig, connectedConfig, err) + retryInst, retryConfig, retryErr := a.retryConnectAfterMySQLMaxUserConnections(effectiveConfig, connectedConfig, err, flight) if retryErr != nil { failedKey := getCacheKey(retryConfig) - a.recordConnectFailureByKey(failedKey, retryErr) - return nil, retryErr + if flightErr := a.recordConnectFailureForFlight(flight, failedKey, retryErr); flightErr != nil { + return databaseConnectResult{}, flightErr + } + return databaseConnectResult{}, retryErr } dbInst = retryInst connectedConfig = retryConfig } - a.clearConnectFailureByKey(initialKey) effectiveConfig = connectedConfig key = getCacheKey(effectiveConfig) shortKey = shortenCacheKey(key) - a.clearConnectFailureByKey(key) now := time.Now() keepAlivePolicy := resolveConnectionKeepAlivePolicy(effectiveConfig) a.mu.Lock() + // A successful driver Connect is not publishable until its flight token is + // revalidated under the cache lock. Close any invalidated instance outside it. + flightErr = a.databaseConnectFlightErrorLocked(flight) + if flightErr == nil { + flight.cacheKey = key + flight.releaseMatchKey = getConnectionReleaseMatchKey(effectiveConfig) + } + if flightErr != nil { + a.mu.Unlock() + _ = dbInst.Close() + return databaseConnectResult{}, flightErr + } if existing, exists = a.dbCache[key]; exists && existing.inst != nil { existing = keepAlivePolicy.apply(existing, now) a.dbCache[key] = existing + if clearErr := a.clearConnectFailuresForFlightLocked(flight, initialKey, key); clearErr != nil { + a.mu.Unlock() + _ = dbInst.Close() + return databaseConnectResult{}, clearErr + } a.mu.Unlock() // Prefer existing cached connection to avoid cache racing duplicates. _ = dbInst.Close() if isFileDB { logger.Infof("并发创建命中已存在文件库连接,关闭新建连接并复用缓存:类型=%s 缓存Key=%s", strings.TrimSpace(effectiveConfig.Type), shortKey) } - return existing.inst, nil + return databaseConnectResult{inst: existing.inst, cacheKey: key}, nil } a.dbCache[key] = cachedDatabase{ inst: dbInst, @@ -1131,10 +1378,16 @@ func (a *App) connectAndCacheDatabase(effectiveConfig connection.ConnectionConfi keepAliveDBType: keepAlivePolicy.dbType, keepAliveRevision: 1, } + if clearErr := a.clearConnectFailuresForFlightLocked(flight, initialKey, key); clearErr != nil { + delete(a.dbCache, key) + a.mu.Unlock() + _ = dbInst.Close() + return databaseConnectResult{}, clearErr + } a.mu.Unlock() logger.Infof("数据库连接成功并写入缓存:%s 缓存Key=%s", formatConnSummary(effectiveConfig), shortKey) - return dbInst, nil + return databaseConnectResult{inst: dbInst, cacheKey: key}, nil } func (a *App) cachedConnectFailureError(effectiveConfig connection.ConnectionConfig, key string, messageKey string) error { @@ -1151,12 +1404,16 @@ func (a *App) cachedConnectFailureError(effectiveConfig connection.ConnectionCon return withLogHint{err: fmt.Errorf("%s", message), logPath: logger.Path()} } -func (a *App) retryConnectAfterMySQLMaxUserConnections(rawConfig connection.ConnectionConfig, failedConfig connection.ConnectionConfig, err error) (db.Database, connection.ConnectionConfig, error) { +func (a *App) retryConnectAfterMySQLMaxUserConnections(rawConfig connection.ConnectionConfig, failedConfig connection.ConnectionConfig, err error, flight *databaseConnectFlight) (db.Database, connection.ConnectionConfig, error) { if !isMySQLMaxUserConnectionsError(err) { return nil, failedConfig, err } - released := a.releaseCachedDatabaseConnectionsForConfig(failedConfig) + excludedFlightID := uint64(0) + if flight != nil { + excludedFlightID = flight.id + } + released := a.releaseCachedDatabaseConnectionsForConfigExcludingFlight(failedConfig, excludedFlightID) logger.Warnf("检测到 MySQL 用户连接数超限,已释放同实例缓存连接:%s 数量=%d", formatConnSummary(failedConfig), released) if released <= 0 { return nil, failedConfig, withMySQLMaxUserConnectionsHint(err, released) @@ -1178,28 +1435,38 @@ func (a *App) getCachedConnectFailureByKey(key string) (cachedConnectFailure, ti return cachedConnectFailure{}, 0, false } - a.mu.RLock() + a.mu.Lock() + defer a.mu.Unlock() entry, exists := a.connectFailures[key] - a.mu.RUnlock() if !exists || entry.err == nil || entry.occurredAt.IsZero() { return cachedConnectFailure{}, 0, false } remaining := dbConnectFailureCooldown - time.Since(entry.occurredAt) if remaining <= 0 { - a.clearConnectFailureByKey(key) + a.clearConnectFailureByKeyLocked(key) return cachedConnectFailure{}, 0, false } return entry, remaining, true } -func (a *App) recordConnectFailureByKey(key string, err error) { - if a == nil || strings.TrimSpace(key) == "" || err == nil { - return +func (a *App) recordConnectFailureForFlight(flight *databaseConnectFlight, key string, err error) error { + if a == nil { + return errDatabaseConnectionShutdown + } + if strings.TrimSpace(key) == "" || err == nil { + return nil } a.mu.Lock() + defer a.mu.Unlock() + if flightErr := a.databaseConnectFlightErrorLocked(flight); flightErr != nil { + return flightErr + } + // Keep the final failure key on the active token so a release that wins + // immediately after this commit can clear the just-recorded cooldown. + flight.cacheKey = key if a.connectFailures == nil { a.connectFailures = make(map[string]cachedConnectFailure) } @@ -1207,19 +1474,24 @@ func (a *App) recordConnectFailureByKey(key string, err error) { occurredAt: time.Now(), err: err, } - a.mu.Unlock() + return nil } -func (a *App) clearConnectFailureByKey(key string) { - if a == nil || strings.TrimSpace(key) == "" { +func (a *App) clearConnectFailureByKeyLocked(key string) { + if strings.TrimSpace(key) == "" || a.connectFailures == nil { return } + delete(a.connectFailures, key) +} - a.mu.Lock() - if a.connectFailures != nil { - delete(a.connectFailures, key) +func (a *App) clearConnectFailuresForFlightLocked(flight *databaseConnectFlight, keys ...string) error { + if flightErr := a.databaseConnectFlightErrorLocked(flight); flightErr != nil { + return flightErr } - a.mu.Unlock() + for _, key := range keys { + a.clearConnectFailureByKeyLocked(key) + } + return nil } func formatConnectFailureCooldown(remaining time.Duration) time.Duration { diff --git a/internal/app/app_db_cache_concurrency_test.go b/internal/app/app_db_cache_concurrency_test.go index dc46f07e..407fb617 100644 --- a/internal/app/app_db_cache_concurrency_test.go +++ b/internal/app/app_db_cache_concurrency_test.go @@ -355,3 +355,660 @@ func TestInvalidateCachedDatabase_SlowCloseDoesNotBlockUnrelatedColdConnect(t *t t.Fatal("expected the slow cached database to be invalidated") } } + +func TestDBReleaseConnection_InvalidatesInflightLeaderAndFollower(t *testing.T) { + installDatabaseCacheConcurrencyTestHooks(t) + + previousMaxProcs := runtime.GOMAXPROCS(1) + t.Cleanup(func() { runtime.GOMAXPROCS(previousMaxProcs) }) + + connectStarted := make(chan struct{}) + releaseConnect := make(chan struct{}) + var connectStartedOnce sync.Once + var closeCalls atomic.Int32 + created := &cacheConcurrencyDB{ + connect: func(connection.ConnectionConfig) error { + connectStartedOnce.Do(func() { close(connectStarted) }) + <-releaseConnect + return nil + }, + close: func() error { + closeCalls.Add(1) + return nil + }, + } + newDatabaseFunc = func(string) (db.Database, error) { return created, nil } + + app := newDatabaseCacheConcurrencyTestApp() + config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "release-race"} + type connectResult struct { + instance db.Database + err error + } + results := make(chan connectResult, 2) + go func() { + instance, err := app.getDatabaseWithPing(config, false) + results <- connectResult{instance: instance, err: err} + }() + select { + case <-connectStarted: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for the leader connection") + } + + followerStarted := make(chan struct{}) + go func() { + close(followerStarted) + instance, err := app.getDatabaseWithPing(config, false) + results <- connectResult{instance: instance, err: err} + }() + <-followerStarted + // Give the follower a scheduling turn to join the blocked singleflight. + // This mirrors the existing high-contention coalescing regression above. + time.Sleep(25 * time.Millisecond) + + releaseResult := app.DBReleaseConnection(config) + if !releaseResult.Success { + close(releaseConnect) + for range 2 { + <-results + } + t.Fatalf("DBReleaseConnection failed: %s", releaseResult.Message) + } + close(releaseConnect) + + for range 2 { + result := <-results + if !errors.Is(result.err, errDatabaseConnectionReleased) { + t.Fatalf("expected release error from the invalidated flight, got %v", result.err) + } + if result.instance != nil { + t.Fatal("an invalidated connection was returned to a caller") + } + } + if got := closeCalls.Load(); got != 1 { + t.Fatalf("expected the late physical connection to be closed once, got %d", got) + } + app.mu.RLock() + cacheSize := len(app.dbCache) + flightCount := len(app.dbConnectFlights) + app.mu.RUnlock() + if cacheSize != 0 { + t.Fatalf("released in-flight connection was resurrected in cache: %d entries", cacheSize) + } + if flightCount != 0 { + t.Fatalf("completed release left %d active flight tokens", flightCount) + } +} + +func TestDBReleaseConnection_ReleaseWinsOverLateConnectFailureWithoutCooldown(t *testing.T) { + installDatabaseCacheConcurrencyTestHooks(t) + + connectStarted := make(chan struct{}) + releaseConnect := make(chan struct{}) + var closeCalls atomic.Int32 + newDatabaseFunc = func(string) (db.Database, error) { + return &cacheConcurrencyDB{ + connect: func(connection.ConnectionConfig) error { + close(connectStarted) + <-releaseConnect + return errors.New("late dial failure") + }, + close: func() error { + closeCalls.Add(1) + return nil + }, + }, nil + } + + app := newDatabaseCacheConcurrencyTestApp() + config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "release-failure-race"} + connectDone := make(chan error, 1) + go func() { + _, err := app.getDatabaseWithPing(config, false) + connectDone <- err + }() + select { + case <-connectStarted: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for the failing connection") + } + if result := app.DBReleaseConnection(config); !result.Success { + close(releaseConnect) + <-connectDone + t.Fatalf("DBReleaseConnection failed: %s", result.Message) + } + close(releaseConnect) + if err := <-connectDone; !errors.Is(err, errDatabaseConnectionReleased) { + t.Fatalf("expected release to win over the late dial error, got %v", err) + } + if got := closeCalls.Load(); got != 1 { + t.Fatalf("expected failed physical connection to close once, got %d", got) + } + app.mu.RLock() + failureCount := len(app.connectFailures) + app.mu.RUnlock() + if failureCount != 0 { + t.Fatalf("released flight poisoned connect cooldown with %d entries", failureCount) + } +} + +func TestDBReleaseConnection_AllowsFreshFileDatabaseFlightBeforeOldConnectFinishes(t *testing.T) { + installDatabaseCacheConcurrencyTestHooks(t) + + firstConnectStarted := make(chan struct{}) + releaseFirstConnect := make(chan struct{}) + secondConnectStarted := make(chan struct{}) + var factoryCalls atomic.Int32 + var firstCloseCalls atomic.Int32 + firstDB := &cacheConcurrencyDB{ + connect: func(connection.ConnectionConfig) error { + close(firstConnectStarted) + <-releaseFirstConnect + return nil + }, + close: func() error { + firstCloseCalls.Add(1) + return nil + }, + } + secondDB := &cacheConcurrencyDB{ + connect: func(connection.ConnectionConfig) error { + close(secondConnectStarted) + return nil + }, + } + newDatabaseFunc = func(string) (db.Database, error) { + if factoryCalls.Add(1) == 1 { + return firstDB, nil + } + return secondDB, nil + } + + app := newDatabaseCacheConcurrencyTestApp() + config := connection.ConnectionConfig{Type: "sqlite", Host: t.TempDir() + "/release-race.sqlite"} + firstDone := make(chan error, 1) + go func() { + _, err := app.getDatabaseWithPing(config, false) + firstDone <- err + }() + select { + case <-firstConnectStarted: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for the old file database flight") + } + + if result := app.DBReleaseConnection(config); !result.Success { + close(releaseFirstConnect) + <-firstDone + t.Fatalf("DBReleaseConnection failed: %s", result.Message) + } + secondDone := make(chan struct { + instance db.Database + err error + }, 1) + go func() { + instance, err := app.getDatabaseWithPing(config, false) + secondDone <- struct { + instance db.Database + err error + }{instance: instance, err: err} + }() + select { + case <-secondConnectStarted: + case <-time.After(2 * time.Second): + close(releaseFirstConnect) + <-firstDone + <-secondDone + t.Fatal("fresh connect joined the released singleflight instead of starting immediately") + } + secondResult := <-secondDone + if secondResult.err != nil || secondResult.instance != secondDB { + close(releaseFirstConnect) + <-firstDone + t.Fatalf("fresh connect failed: instance=%p err=%v", secondResult.instance, secondResult.err) + } + + close(releaseFirstConnect) + if err := <-firstDone; !errors.Is(err, errDatabaseConnectionReleased) { + t.Fatalf("expected old flight to stay invalidated, got %v", err) + } + if got := firstCloseCalls.Load(); got != 1 { + t.Fatalf("expected old file database instance to close once, got %d", got) + } + key := getCacheKey(config) + app.mu.RLock() + entry := app.dbCache[key] + cacheSize := len(app.dbCache) + app.mu.RUnlock() + if cacheSize != 1 || entry.inst != secondDB { + t.Fatalf("old flight displaced the fresh cache entry: size=%d instance=%p", cacheSize, entry.inst) + } +} + +func TestDBReleaseConnection_InvalidatesAllInflightDatabaseKeysForSameInstance(t *testing.T) { + installDatabaseCacheConcurrencyTestHooks(t) + + connectStarted := make(chan string, 2) + releaseConnect := make(chan struct{}) + var closeCalls atomic.Int32 + newDatabaseFunc = func(string) (db.Database, error) { + return &cacheConcurrencyDB{ + connect: func(config connection.ConnectionConfig) error { + connectStarted <- config.Database + <-releaseConnect + return nil + }, + close: func() error { + closeCalls.Add(1) + return nil + }, + }, nil + } + + app := newDatabaseCacheConcurrencyTestApp() + configs := []connection.ConnectionConfig{ + {Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "main"}, + {Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "analytics", ConnectionParams: "application_name=gonavi"}, + } + errorsByDatabase := make(chan error, len(configs)) + for _, config := range configs { + config := config + go func() { + _, err := app.getDatabaseWithPing(config, false) + errorsByDatabase <- err + }() + } + started := make(map[string]bool, len(configs)) + for len(started) < len(configs) { + select { + case databaseName := <-connectStarted: + started[databaseName] = true + case <-time.After(2 * time.Second): + close(releaseConnect) + for range configs { + <-errorsByDatabase + } + t.Fatal("timed out waiting for both same-instance flights") + } + } + if result := app.DBReleaseConnection(configs[0]); !result.Success { + close(releaseConnect) + for range configs { + <-errorsByDatabase + } + t.Fatalf("DBReleaseConnection failed: %s", result.Message) + } + close(releaseConnect) + for range configs { + if err := <-errorsByDatabase; !errors.Is(err, errDatabaseConnectionReleased) { + t.Fatalf("expected every same-instance flight to be invalidated, got %v", err) + } + } + if got := closeCalls.Load(); got != int32(len(configs)) { + t.Fatalf("expected every late connection to close, got %d", got) + } + app.mu.RLock() + cacheSize := len(app.dbCache) + app.mu.RUnlock() + if cacheSize != 0 { + t.Fatalf("same-instance release left %d cache entries", cacheSize) + } +} + +func TestBeginDatabaseShutdown_InvalidatesInflightConnectAndRejectsNewConnects(t *testing.T) { + installDatabaseCacheConcurrencyTestHooks(t) + + connectStarted := make(chan struct{}) + releaseConnect := make(chan struct{}) + var factoryCalls atomic.Int32 + var closeCalls atomic.Int32 + newDatabaseFunc = func(string) (db.Database, error) { + factoryCalls.Add(1) + return &cacheConcurrencyDB{ + connect: func(connection.ConnectionConfig) error { + close(connectStarted) + <-releaseConnect + return nil + }, + close: func() error { + closeCalls.Add(1) + return nil + }, + }, nil + } + + app := newDatabaseCacheConcurrencyTestApp() + config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "shutdown-race"} + connectDone := make(chan error, 1) + go func() { + _, err := app.getDatabaseWithPing(config, false) + connectDone <- err + }() + select { + case <-connectStarted: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for the in-flight connection") + } + + app.beginDatabaseShutdown() + close(releaseConnect) + if err := <-connectDone; !errors.Is(err, errDatabaseConnectionShutdown) { + t.Fatalf("expected shutdown error from the late flight, got %v", err) + } + if got := closeCalls.Load(); got != 1 { + t.Fatalf("expected late shutdown connection to close once, got %d", got) + } + if _, err := app.getDatabase(config); !errors.Is(err, errDatabaseConnectionShutdown) { + t.Fatalf("expected new connect after shutdown to be rejected, got %v", err) + } + if got := factoryCalls.Load(); got != 1 { + t.Fatalf("shutdown created an additional physical connection: %d", got) + } + app.mu.RLock() + cacheSize := len(app.dbCache) + app.mu.RUnlock() + if cacheSize != 0 { + t.Fatalf("shutdown race resurrected %d cache entries", cacheSize) + } +} + +func TestCloseCachedDatabasesForShutdown_SlowCloseDoesNotHoldCacheLock(t *testing.T) { + closeStarted := make(chan struct{}) + releaseClose := make(chan struct{}) + app := newDatabaseCacheConcurrencyTestApp() + config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "shutdown-close"} + app.dbCache[getCacheKey(config)] = cachedDatabase{ + inst: &cacheConcurrencyDB{close: func() error { + close(closeStarted) + <-releaseClose + return nil + }}, + config: normalizeCacheKeyConfig(config), + } + + closeDone := make(chan struct{}) + go func() { + app.closeCachedDatabasesForShutdown() + close(closeDone) + }() + select { + case <-closeStarted: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for the slow shutdown close") + } + + lockAcquired := make(chan struct{}) + go func() { + app.mu.Lock() + app.mu.Unlock() + close(lockAcquired) + }() + select { + case <-lockAcquired: + case <-time.After(time.Second): + close(releaseClose) + <-closeDone + t.Fatal("slow shutdown Close held the global database cache lock") + } + close(releaseClose) + select { + case <-closeDone: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for shutdown close to finish") + } +} + +func TestRecordConnectFailureForFlight_CancellationAndCooldownCommitAreLinearized(t *testing.T) { + config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "failure-linearization"} + key := getCacheKey(config) + dialErr := errors.New("dial failed") + + t.Run("release before record", func(t *testing.T) { + app := newDatabaseCacheConcurrencyTestApp() + flight, err := app.beginDatabaseConnectFlight(key, config) + if err != nil { + t.Fatalf("begin flight: %v", err) + } + defer app.finishDatabaseConnectFlight(flight) + + // Reproduce the old check-then-record window: validation succeeds, then a + // release wins before the cooldown write is attempted. + if err := app.databaseConnectFlightError(flight); err != nil { + t.Fatalf("flight unexpectedly invalid before release: %v", err) + } + app.releaseCachedDatabaseConnectionsForConfig(config) + if err := app.recordConnectFailureForFlight(flight, key, dialErr); !errors.Is(err, errDatabaseConnectionReleased) { + t.Fatalf("expected released flight to reject cooldown write, got %v", err) + } + app.mu.RLock() + failureCount := len(app.connectFailures) + app.mu.RUnlock() + if failureCount != 0 { + t.Fatalf("released flight wrote %d stale cooldown entries", failureCount) + } + }) + + t.Run("release after record", func(t *testing.T) { + app := newDatabaseCacheConcurrencyTestApp() + flight, err := app.beginDatabaseConnectFlight(key, config) + if err != nil { + t.Fatalf("begin flight: %v", err) + } + defer app.finishDatabaseConnectFlight(flight) + + if err := app.recordConnectFailureForFlight(flight, key, dialErr); err != nil { + t.Fatalf("record active flight failure: %v", err) + } + app.releaseCachedDatabaseConnectionsForConfig(config) + app.mu.RLock() + _, exists := app.connectFailures[key] + app.mu.RUnlock() + if exists { + t.Fatal("release did not clear cooldown committed by its active flight") + } + }) + + t.Run("shutdown before record", func(t *testing.T) { + app := newDatabaseCacheConcurrencyTestApp() + flight, err := app.beginDatabaseConnectFlight(key, config) + if err != nil { + t.Fatalf("begin flight: %v", err) + } + defer app.finishDatabaseConnectFlight(flight) + + if err := app.databaseConnectFlightError(flight); err != nil { + t.Fatalf("flight unexpectedly invalid before shutdown: %v", err) + } + app.beginDatabaseShutdown() + if err := app.recordConnectFailureForFlight(flight, key, dialErr); !errors.Is(err, errDatabaseConnectionShutdown) { + t.Fatalf("expected shutdown to reject cooldown write, got %v", err) + } + app.mu.RLock() + failureCount := len(app.connectFailures) + app.mu.RUnlock() + if failureCount != 0 { + t.Fatalf("shutdown flight wrote %d stale cooldown entries", failureCount) + } + }) +} + +func TestClearConnectFailuresForFlightLocked_StaleFlightCannotEraseFreshCooldown(t *testing.T) { + app := newDatabaseCacheConcurrencyTestApp() + config := connection.ConnectionConfig{Type: "postgres", Host: "127.0.0.1", Port: 5432, User: "postgres", Database: "clear-linearization"} + key := getCacheKey(config) + + staleFlight, err := app.beginDatabaseConnectFlight(key, config) + if err != nil { + t.Fatalf("begin stale flight: %v", err) + } + defer app.finishDatabaseConnectFlight(staleFlight) + if err := app.databaseConnectFlightError(staleFlight); err != nil { + t.Fatalf("stale flight unexpectedly invalid before release: %v", err) + } + app.releaseCachedDatabaseConnectionsForConfig(config) + + freshFlight, err := app.beginDatabaseConnectFlight(key, config) + if err != nil { + t.Fatalf("begin fresh flight: %v", err) + } + defer app.finishDatabaseConnectFlight(freshFlight) + freshErr := errors.New("fresh dial failure") + if err := app.recordConnectFailureForFlight(freshFlight, key, freshErr); err != nil { + t.Fatalf("record fresh cooldown: %v", err) + } + + // Reproduce the old validation-then-clear window. The stale flight must be + // revalidated in the same critical section as the deletion. + app.mu.Lock() + clearErr := app.clearConnectFailuresForFlightLocked(staleFlight, key) + failure, exists := app.connectFailures[key] + app.mu.Unlock() + if !errors.Is(clearErr, errDatabaseConnectionReleased) { + t.Fatalf("expected stale clear to be rejected, got %v", clearErr) + } + if !exists || !errors.Is(failure.err, freshErr) { + t.Fatalf("stale flight erased or replaced fresh cooldown: exists=%t err=%v", exists, failure.err) + } +} + +func TestDBReleaseConnection_ForgetIsLinearizedBeforeFreshSingleflight(t *testing.T) { + installDatabaseCacheConcurrencyTestHooks(t) + + firstConnectStarted := make(chan struct{}) + releaseFirstConnect := make(chan struct{}) + secondConnectStarted := make(chan struct{}) + releaseSecondConnect := make(chan struct{}) + duplicateConnectStarted := make(chan struct{}) + var duplicateConnectOnce sync.Once + var factoryCalls atomic.Int32 + firstDB := &cacheConcurrencyDB{connect: func(connection.ConnectionConfig) error { + close(firstConnectStarted) + <-releaseFirstConnect + return nil + }} + secondDB := &cacheConcurrencyDB{connect: func(connection.ConnectionConfig) error { + close(secondConnectStarted) + <-releaseSecondConnect + return nil + }} + newDatabaseFunc = func(string) (db.Database, error) { + switch factoryCalls.Add(1) { + case 1: + return firstDB, nil + case 2: + return secondDB, nil + default: + return &cacheConcurrencyDB{connect: func(connection.ConnectionConfig) error { + duplicateConnectOnce.Do(func() { close(duplicateConnectStarted) }) + return nil + }}, nil + } + } + + app := newDatabaseCacheConcurrencyTestApp() + forgetHeldCacheLock := make(chan bool, 1) + app.dbConnectBeforeForgetHook = func() { + if app.mu.TryLock() { + app.mu.Unlock() + forgetHeldCacheLock <- false + return + } + forgetHeldCacheLock <- true + } + config := connection.ConnectionConfig{Type: "sqlite", Host: t.TempDir() + "/forget-aba.sqlite"} + + firstDone := make(chan error, 1) + go func() { + _, err := app.getDatabaseWithPing(config, false) + firstDone <- err + }() + select { + case <-firstConnectStarted: + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for old singleflight") + } + if result := app.DBReleaseConnection(config); !result.Success { + close(releaseFirstConnect) + <-firstDone + t.Fatalf("DBReleaseConnection failed: %s", result.Message) + } + if held := <-forgetHeldCacheLock; !held { + close(releaseFirstConnect) + <-firstDone + t.Fatal("singleflight Forget ran after releasing the cache lock, reopening the ABA window") + } + + secondDone := make(chan struct { + instance db.Database + err error + }, 1) + go func() { + instance, err := app.getDatabaseWithPing(config, false) + secondDone <- struct { + instance db.Database + err error + }{instance: instance, err: err} + }() + select { + case <-secondConnectStarted: + case <-time.After(2 * time.Second): + close(releaseFirstConnect) + <-firstDone + t.Fatal("timed out waiting for fresh singleflight") + } + + // Let the forgotten old group finish only after the fresh group exists. It + // must neither remove nor split the fresh group. + close(releaseFirstConnect) + if err := <-firstDone; !errors.Is(err, errDatabaseConnectionReleased) { + close(releaseSecondConnect) + <-secondDone + t.Fatalf("expected old flight to remain released, got %v", err) + } + + const followers = 16 + results := make(chan struct { + instance db.Database + err error + }, followers) + start := make(chan struct{}) + for range followers { + go func() { + <-start + instance, err := app.getDatabaseWithPing(config, false) + results <- struct { + instance db.Database + err error + }{instance: instance, err: err} + }() + } + close(start) + for range 32 { + runtime.Gosched() + } + select { + case <-duplicateConnectStarted: + close(releaseSecondConnect) + <-secondDone + for range followers { + <-results + } + t.Fatal("fresh singleflight was forgotten and split into duplicate physical connects") + default: + } + close(releaseSecondConnect) + secondResult := <-secondDone + if secondResult.err != nil || secondResult.instance != secondDB { + t.Fatalf("fresh leader failed: instance=%p err=%v", secondResult.instance, secondResult.err) + } + for range followers { + result := <-results + if result.err != nil || result.instance != secondDB { + t.Fatalf("fresh follower failed: instance=%p err=%v", result.instance, result.err) + } + } + if got := factoryCalls.Load(); got != 2 { + t.Fatalf("expected exactly two physical flights across release, got %d", got) + } +}