Move auth queries to primary, redo health check loop

This commit is contained in:
binwiederhier
2026-03-11 20:26:29 -04:00
parent ab33ac7ae5
commit ac65df1e83
2 changed files with 49 additions and 35 deletions
+42 -28
View File
@@ -1,6 +1,7 @@
package db package db
import ( import (
"context"
"database/sql" "database/sql"
"sync/atomic" "sync/atomic"
"time" "time"
@@ -9,7 +10,8 @@ import (
) )
const ( const (
replicaHealthCheckInterval = 5 * time.Second replicaHealthCheckInterval = 30 * time.Second
replicaHealthCheckTimeout = 2 * time.Second
) )
// Beginner is an interface for types that can begin a database transaction. // Beginner is an interface for types that can begin a database transaction.
@@ -25,25 +27,29 @@ type DB struct {
primary *sql.DB primary *sql.DB
replicas []*replica replicas []*replica
counter atomic.Uint64 counter atomic.Uint64
cancel context.CancelFunc
} }
type replica struct { type replica struct {
db *sql.DB db *sql.DB
healthy atomic.Bool healthy atomic.Bool
lastChecked atomic.Int64
} }
// NewDB creates a new DB that wraps the given primary and optional replica connections. // NewDB creates a new DB that wraps the given primary and optional replica connections.
// If replicas is nil or empty, ReadOnly() simply returns the primary. // If replicas is nil or empty, ReadOnly() simply returns the primary.
// Replicas start unhealthy and are checked immediately by a background goroutine.
func NewDB(primary *sql.DB, replicas []*sql.DB) *DB { func NewDB(primary *sql.DB, replicas []*sql.DB) *DB {
ctx, cancel := context.WithCancel(context.Background())
d := &DB{ d := &DB{
primary: primary, primary: primary,
replicas: make([]*replica, len(replicas)), replicas: make([]*replica, len(replicas)),
cancel: cancel,
} }
for i, r := range replicas { for i, r := range replicas {
rep := &replica{db: r} d.replicas[i] = &replica{db: r} // healthy defaults to false
rep.healthy.Store(true) }
d.replicas[i] = rep if len(d.replicas) > 0 {
go d.healthCheckLoop(ctx)
} }
return d return d
} }
@@ -79,8 +85,9 @@ func (d *DB) Ping() error {
return d.primary.Ping() return d.primary.Ping()
} }
// Close closes the primary database and all replicas. // Close closes the primary database and all replicas, and stops the health-check goroutine.
func (d *DB) Close() error { func (d *DB) Close() error {
d.cancel()
for _, r := range d.replicas { for _, r := range d.replicas {
r.db.Close() r.db.Close()
} }
@@ -88,9 +95,7 @@ func (d *DB) Close() error {
} }
// ReadOnly returns a *sql.DB suitable for read-only queries. It round-robins across healthy // ReadOnly returns a *sql.DB suitable for read-only queries. It round-robins across healthy
// replicas. If a replica's health status is stale (older than replicaHealthCheckInterval), it // replicas. If all replicas are unhealthy or none are configured, the primary is returned.
// is re-checked with a ping. If all replicas are unhealthy or none are configured, the primary
// is returned.
func (d *DB) ReadOnly() *sql.DB { func (d *DB) ReadOnly() *sql.DB {
if len(d.replicas) == 0 { if len(d.replicas) == 0 {
return d.primary return d.primary
@@ -99,34 +104,43 @@ func (d *DB) ReadOnly() *sql.DB {
start := int(d.counter.Add(1) - 1) start := int(d.counter.Add(1) - 1)
for i := 0; i < n; i++ { for i := 0; i < n; i++ {
r := d.replicas[(start+i)%n] r := d.replicas[(start+i)%n]
if d.isHealthy(r) { if r.healthy.Load() {
return r.db return r.db
} }
} }
return d.primary return d.primary
} }
// isHealthy returns whether the replica is healthy. If the cached health status is stale, // healthCheckLoop checks replicas immediately, then periodically on a ticker.
// it pings the replica and updates the cache. func (d *DB) healthCheckLoop(ctx context.Context) {
func (d *DB) isHealthy(r *replica) bool { d.checkReplicas(ctx)
now := time.Now().Unix() for {
lastChecked := r.lastChecked.Load() select {
if now-lastChecked >= int64(replicaHealthCheckInterval.Seconds()) { case <-ctx.Done():
if r.lastChecked.CompareAndSwap(lastChecked, now) { return
wasHealthy := r.healthy.Load() case <-time.After(replicaHealthCheckInterval):
if err := r.db.Ping(); err != nil { d.checkReplicas(ctx)
r.healthy.Store(false) }
if wasHealthy { }
log.Error("Database replica is now unhealthy: %s", err) }
}
return false // checkReplicas pings each replica with a timeout and updates its health status.
func (d *DB) checkReplicas(ctx context.Context) {
for _, r := range d.replicas {
wasHealthy := r.healthy.Load()
pingCtx, cancel := context.WithTimeout(ctx, replicaHealthCheckTimeout)
err := r.db.PingContext(pingCtx)
cancel()
if err != nil {
r.healthy.Store(false)
if wasHealthy {
log.Error("Database replica is now unhealthy: %s", err)
} }
} else {
r.healthy.Store(true) r.healthy.Store(true)
if !wasHealthy { if !wasHealthy {
log.Info("Database replica is now healthy again") log.Info("Database replica is now healthy again")
} }
return true
} }
} }
return r.healthy.Load()
} }
+7 -7
View File
@@ -388,7 +388,7 @@ func (a *Manager) writeUserStatsQueue() error {
// User returns the user with the given username if it exists, or ErrUserNotFound otherwise // User returns the user with the given username if it exists, or ErrUserNotFound otherwise
func (a *Manager) User(username string) (*User, error) { func (a *Manager) User(username string) (*User, error) {
rows, err := a.db.ReadOnly().Query(a.queries.selectUserByName, username) rows, err := a.db.Query(a.queries.selectUserByName, username)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -397,7 +397,7 @@ func (a *Manager) User(username string) (*User, error) {
// UserByID returns the user with the given ID if it exists, or ErrUserNotFound otherwise // UserByID returns the user with the given ID if it exists, or ErrUserNotFound otherwise
func (a *Manager) UserByID(id string) (*User, error) { func (a *Manager) UserByID(id string) (*User, error) {
rows, err := a.db.ReadOnly().Query(a.queries.selectUserByID, id) rows, err := a.db.Query(a.queries.selectUserByID, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -406,7 +406,7 @@ func (a *Manager) UserByID(id string) (*User, error) {
// userByToken returns the user with the given token if it exists and is not expired, or ErrUserNotFound otherwise // userByToken returns the user with the given token if it exists and is not expired, or ErrUserNotFound otherwise
func (a *Manager) userByToken(token string) (*User, error) { func (a *Manager) userByToken(token string) (*User, error) {
rows, err := a.db.ReadOnly().Query(a.queries.selectUserByToken, token, time.Now().Unix()) rows, err := a.db.Query(a.queries.selectUserByToken, token, time.Now().Unix())
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -642,7 +642,7 @@ func (a *Manager) AllowReservation(username string, topic string) error {
// - Furthermore, the query prioritizes more specific permissions (longer!) over more generic ones, e.g. "test*" > "*" // - Furthermore, the query prioritizes more specific permissions (longer!) over more generic ones, e.g. "test*" > "*"
// - It also prioritizes write permissions over read permissions // - It also prioritizes write permissions over read permissions
func (a *Manager) authorizeTopicAccess(usernameOrEveryone, topic string) (read, write, found bool, err error) { func (a *Manager) authorizeTopicAccess(usernameOrEveryone, topic string) (read, write, found bool, err error) {
rows, err := a.db.ReadOnly().Query(a.queries.selectTopicPerms, Everyone, usernameOrEveryone, topic) rows, err := a.db.Query(a.queries.selectTopicPerms, Everyone, usernameOrEveryone, topic)
if err != nil { if err != nil {
return false, false, false, err return false, false, false, err
} }
@@ -779,7 +779,7 @@ func (a *Manager) Reservations(username string) ([]Reservation, error) {
// HasReservation returns true if the given topic access is owned by the user // HasReservation returns true if the given topic access is owned by the user
func (a *Manager) HasReservation(username, topic string) (bool, error) { func (a *Manager) HasReservation(username, topic string) (bool, error) {
rows, err := a.db.ReadOnly().Query(a.queries.selectUserHasReservation, username, escapeUnderscore(topic)) rows, err := a.db.Query(a.queries.selectUserHasReservation, username, escapeUnderscore(topic))
if err != nil { if err != nil {
return false, err return false, err
} }
@@ -813,7 +813,7 @@ func (a *Manager) ReservationsCount(username string) (int64, error) {
// ReservationOwner returns user ID of the user that owns this topic, or an empty string if it's not owned by anyone // ReservationOwner returns user ID of the user that owns this topic, or an empty string if it's not owned by anyone
func (a *Manager) ReservationOwner(topic string) (string, error) { func (a *Manager) ReservationOwner(topic string) (string, error) {
rows, err := a.db.ReadOnly().Query(a.queries.selectUserReservationsOwner, escapeUnderscore(topic)) rows, err := a.db.Query(a.queries.selectUserReservationsOwner, escapeUnderscore(topic))
if err != nil { if err != nil {
return "", err return "", err
} }
@@ -830,7 +830,7 @@ func (a *Manager) ReservationOwner(topic string) (string, error) {
// otherAccessCount returns the number of access entries for the given topic that are not owned by the user // otherAccessCount returns the number of access entries for the given topic that are not owned by the user
func (a *Manager) otherAccessCount(username, topic string) (int, error) { func (a *Manager) otherAccessCount(username, topic string) (int, error) {
rows, err := a.db.ReadOnly().Query(a.queries.selectOtherAccessCount, escapeUnderscore(topic), escapeUnderscore(topic), username) rows, err := a.db.Query(a.queries.selectOtherAccessCount, escapeUnderscore(topic), escapeUnderscore(topic), username)
if err != nil { if err != nil {
return 0, err return 0, err
} }