mirror of
https://github.com/multipleof4/ntfy.git
synced 2026-10-08 21:05:21 +00:00
Move auth queries to primary, redo health check loop
This commit is contained in:
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user