diff --git a/errors.go b/errors.go index e6b16011..b140efbe 100644 --- a/errors.go +++ b/errors.go @@ -49,4 +49,5 @@ var ( ErrMissingConnURL = errors.New(`upper: missing DSN`) ErrNotImplemented = errors.New(`upper: call not implemented`) ErrAlreadyWithinTransaction = errors.New(`upper: already within a transaction`) + ErrServerRefusedConnection = errors.New(`upper: database server refused connection`) ) diff --git a/internal/sqladapter/database.go b/internal/sqladapter/database.go index 7c1fa4d0..13af8e22 100644 --- a/internal/sqladapter/database.go +++ b/internal/sqladapter/database.go @@ -2,6 +2,9 @@ package sqladapter import ( "database/sql" + "database/sql/driver" + "errors" + "io" "math" "strconv" "sync" @@ -14,13 +17,48 @@ import ( "upper.io/db.v2/lib/sqlbuilder" ) +// A list of errors that mean the server is not working and that we should +// try to connect and retry the query. +var recoverableErrors = []error{ + io.EOF, + driver.ErrBadConn, + db.ErrNotConnected, + db.ErrTooManyClients, + db.ErrServerRefusedConnection, +} + +var ( + // If a query fails with a recoverable error the connection is going to be + // re-estalished and the query can be retried, each retry adds a max wait + // time of maxConnectionRetryTime + maxQueryRetryAttempts = 3 + + // Minimum interval when waiting before trying to reconnect. + minConnectionRetryInterval = time.Millisecond * 100 + + // Maximum interval when waiting before trying to reconnect. + maxConnectionRetryInterval = time.Millisecond * 2500 + + // Maximun time each connection retry attempt can take. + maxConnectionRetryTime = time.Second * 5 + + // Maximun reconnection attempts per session before giving up. + maxReconnectionAttempts uint64 = 12 +) + var ( - lastSessID uint64 - lastTxID uint64 + errNothingToRecoverFrom = errors.New("Nothing to recover from") + errUnableToRecover = errors.New("Unable to recover from this error") ) -// HasCleanUp is implemented by structs that have a clean up routine that needs -// to be called before Close(). +var ( + lastSessID uint64 + lastTxID uint64 + lastOperationID uint64 +) + +// HasCleanUp is implemented by structs that have a clean up routine that +// needs to be called before Close(). type HasCleanUp interface { CleanUp() error } @@ -78,29 +116,41 @@ type BaseDatabase interface { // NewBaseDatabase provides a BaseDatabase given a PartialDatabase func NewBaseDatabase(p PartialDatabase) BaseDatabase { d := &database{ - PartialDatabase: p, + PartialDatabase: p, + cachedCollections: cache.NewCache(), cachedStatements: cache.NewCache(), + connFn: defaultConnFn, } return d } +var defaultConnFn = func() error { + return errors.New("No connection function was defined.") +} + // database is the actual implementation of Database and joins methods from // BaseDatabase and PartialDatabase type database struct { PartialDatabase baseTx BaseTx + connectMu sync.Mutex collectionMu sync.Mutex - databaseMu sync.Mutex - name string - sess *sql.DB - sessMu sync.Mutex + connFn func() error + + name string + + sess *sql.DB + sessErr error + sessMu sync.Mutex sessID uint64 txID uint64 + connectAttempts uint64 + cachedStatements *cache.Cache cachedCollections *cache.Cache @@ -111,8 +161,74 @@ var ( _ = db.Database(&database{}) ) +func (d *database) reconnect() error { + if d.Transaction() != nil { + // Don't even attempt to recover from within a transaction, this is not + // possible. + return errors.New("Can't recover from within a bad transaction.") + } + + err := d.PartialDatabase.Err(d.Ping()) + if err == nil { + return nil + } + + return d.connect(d.connFn) +} + +func (d *database) connect(connFn func() error) error { + if connFn == nil { + return errors.New("Missing connect function") + } + + d.connectMu.Lock() + defer d.connectMu.Unlock() + + // Attempt to (re)connect + if atomic.AddUint64(&d.connectAttempts, 1) >= maxReconnectionAttempts { + return db.ErrGivingUpTryingToConnect + } + + waitTime := minConnectionRetryInterval + + for start, i := time.Now(), 1; time.Now().Sub(start) < maxConnectionRetryTime; i++ { + waitTime = time.Duration(i) * minConnectionRetryInterval + if waitTime > maxConnectionRetryInterval { + waitTime = maxConnectionRetryInterval + } + // Wait a bit until retrying. + if waitTime > time.Duration(0) { + time.Sleep(waitTime) + } + + err := connFn() + if err == nil { + atomic.StoreUint64(&d.connectAttempts, 0) + d.connFn = connFn + return nil + } + + if !d.isRecoverableError(err) { + return err + } + } + + return db.ErrGivingUpTryingToConnect +} + // Session returns the underlying *sql.DB func (d *database) Session() *sql.DB { + if atomic.LoadUint64(&d.sessID) == 0 { + // This means the session is connecting for the first time, in this case we + // don't block because the session hasn't been returned yet. + return d.sess + } + + // Prevents goroutines from using the session until the connection is + // re-established. + d.connectMu.Lock() + defer d.connectMu.Unlock() + return d.sess } @@ -133,55 +249,113 @@ func (d *database) BindTx(t *sql.Tx) error { // Tx returns a BaseTx, which, if not nil, means that this session is within a // transaction func (d *database) Transaction() BaseTx { + if atomic.LoadUint64(&d.sessID) == 0 { + // This means the session is connecting for the first time, in this case we + // don't block because the session hasn't been returned yet. + return d.baseTx + } + + d.sessMu.Lock() + defer d.sessMu.Unlock() + return d.baseTx } // Name returns the database named func (d *database) Name() string { - d.databaseMu.Lock() - defer d.databaseMu.Unlock() + return d.name +} - if d.name == "" { - d.name, _ = d.PartialDatabase.FindDatabaseName() +func (d *database) getDBName() error { + name, err := d.PartialDatabase.FindDatabaseName() + if err != nil { + return err } - - return d.name + d.name = name + return nil } // BindSession binds a *sql.DB into *database func (d *database) BindSession(sess *sql.DB) error { + if err := sess.Ping(); err != nil { + return err + } + d.sessMu.Lock() + if d.sess != nil { + d.ClearCache() + d.sess.Close() // Close before rebind. + } d.sess = sess d.sessMu.Unlock() - if err := d.Ping(); err != nil { + // Does this session already have a session ID? + if atomic.LoadUint64(&d.sessID) != 0 { + return nil + } + + // Is this connection really working? + if err := d.getDBName(); err != nil { return err } + // Assign an ID if everyting was OK. d.sessID = newSessionID() - name, err := d.PartialDatabase.FindDatabaseName() - if err != nil { - return err + return nil +} + +func (d *database) isRecoverableError(err error) bool { + err = d.PartialDatabase.Err(err) + for i := 0; i < len(recoverableErrors); i++ { + if err == recoverableErrors[i] { + return true + } } + return false +} - d.name = name +// recoverFromErr attempts to reestablish a connection after a temporary error, +// returns nil if the connection was reestablished and the query can be retried. +func (d *database) recoverFromErr(err error) error { + if err == nil { + return errNothingToRecoverFrom + } - return nil + if d.isRecoverableError(err) { + err := d.reconnect() + return err + } + + // This is not an error we can recover from. + return errUnableToRecover } // Ping checks whether a connection to the database is still alive by pinging // it func (d *database) Ping() error { - if d.sess != nil { - return d.sess.Ping() + if sess := d.Session(); sess != nil { + if err := sess.Ping(); err != nil { + return err + } + + _, err := sess.Exec("SELECT 1") + if err != nil { + return err + } + return nil } - return nil + if tx := d.Transaction(); tx != nil { + // When upper wraps a transaction with no original session. + return nil + } + return db.ErrNotConnected } // ClearCache removes all caches. func (d *database) ClearCache() { d.collectionMu.Lock() defer d.collectionMu.Unlock() + d.cachedCollections.Clear() d.cachedStatements.Clear() if d.template != nil { @@ -197,7 +371,7 @@ func (d *database) Close() error { d.baseTx = nil d.sessMu.Unlock() }() - if d.sess != nil { + if sess := d.Session(); sess != nil { if cleaner, ok := d.PartialDatabase.(HasCleanUp); ok { cleaner.CleanUp() } @@ -210,6 +384,7 @@ func (d *database) Close() error { return d.sess.Close() } + // Don't close the parent session if within a transaction. if !tx.Committed() { tx.Rollback() } @@ -224,9 +399,9 @@ func (d *database) Collection(name string) db.Collection { h := cache.String(name) - ccol, ok := d.cachedCollections.ReadRaw(h) + cachedCollection, ok := d.cachedCollections.ReadRaw(h) if ok { - return ccol.(db.Collection) + return cachedCollection.(db.Collection) } col := d.PartialDatabase.NewLocalCollection(name) @@ -235,22 +410,39 @@ func (d *database) Collection(name string) db.Collection { return col } +func (d *database) prepareAndExec(stmt *exql.Statement, args ...interface{}) (string, sql.Result, error) { + p, query, err := d.prepareStatement(stmt) + if err != nil { + return query, nil, err + } + + if execer, ok := d.PartialDatabase.(HasStatementExec); ok { + res, err := execer.StatementExec(p.Stmt, args...) + return query, res, err + } + + res, err := p.Exec(args...) + return query, res, err +} + // StatementExec compiles and executes a statement that does not return any // rows. func (d *database) StatementExec(stmt *exql.Statement, args ...interface{}) (res sql.Result, err error) { var query string + queryID := newOperationID() + if db.Conf.LoggingEnabled() { defer func(start time.Time) { - status := db.QueryStatus{ - TxID: d.txID, - SessID: d.sessID, - Query: query, - Args: args, - Err: err, - Start: start, - End: time.Now(), + TxID: d.txID, + SessID: d.sessID, + QueryID: queryID, + Query: query, + Args: args, + Err: err, + Start: start, + End: time.Now(), } if res != nil { @@ -267,47 +459,77 @@ func (d *database) StatementExec(stmt *exql.Statement, args ...interface{}) (res }(time.Now()) } - var p *Stmt - if p, query, err = d.prepareStatement(stmt); err != nil { - return nil, err + for i := 0; ; i++ { + query, res, err = d.prepareAndExec(stmt, args...) + if err == nil || i >= maxQueryRetryAttempts { + return res, err + } + + // Try to recover + if recoverErr := d.recoverFromErr(err); recoverErr != nil { + return nil, err // Unable to recover. + } } - defer p.Close() - if execer, ok := d.PartialDatabase.(HasStatementExec); ok { - res, err = execer.StatementExec(p.Stmt, args...) - return + panic("reached") +} + +func (d *database) prepareAndQuery(stmt *exql.Statement, args ...interface{}) (string, *sql.Rows, error) { + p, query, err := d.prepareStatement(stmt) + if err != nil { + return query, nil, err } - res, err = p.Exec(args...) - return + rows, err := p.Query(args...) + return query, rows, err } // StatementQuery compiles and executes a statement that returns rows. func (d *database) StatementQuery(stmt *exql.Statement, args ...interface{}) (rows *sql.Rows, err error) { var query string + queryID := newOperationID() + if db.Conf.LoggingEnabled() { defer func(start time.Time) { db.Log(&db.QueryStatus{ - TxID: d.txID, - SessID: d.sessID, - Query: query, - Args: args, - Err: err, - Start: start, - End: time.Now(), + TxID: d.txID, + SessID: d.sessID, + QueryID: queryID, + Query: query, + Args: args, + Err: err, + Start: start, + End: time.Now(), }) }(time.Now()) } - var p *Stmt - if p, query, err = d.prepareStatement(stmt); err != nil { - return nil, err + for i := 0; ; i++ { + query, rows, err = d.prepareAndQuery(stmt, args...) + if err == nil || i >= maxQueryRetryAttempts { + return rows, err + } + + // Try to recover + if recoverErr := d.recoverFromErr(err); recoverErr != nil { + return nil, err // Unable to recover. + } + } + + panic("reached") +} + +func (d *database) prepareAndQueryRow(stmt *exql.Statement, args ...interface{}) (string, *sql.Row, error) { + p, query, err := d.prepareStatement(stmt) + if err != nil { + return query, nil, err } - defer p.Close() - rows, err = p.Query(args...) - return + // Would be nice to find a way to check if this succeeded before using + // Scan. + rows, err := p.QueryRow(args...), nil + return query, rows, nil } // StatementQueryRow compiles and executes a statement that returns at most one @@ -315,62 +537,70 @@ func (d *database) StatementQuery(stmt *exql.Statement, args ...interface{}) (ro func (d *database) StatementQueryRow(stmt *exql.Statement, args ...interface{}) (row *sql.Row, err error) { var query string + queryID := newOperationID() + if db.Conf.LoggingEnabled() { defer func(start time.Time) { db.Log(&db.QueryStatus{ - TxID: d.txID, - SessID: d.sessID, - Query: query, - Args: args, - Err: err, - Start: start, - End: time.Now(), + TxID: d.txID, + SessID: d.sessID, + QueryID: queryID, + Query: query, + Args: args, + Err: err, + Start: start, + End: time.Now(), }) }(time.Now()) } - var p *Stmt - if p, query, err = d.prepareStatement(stmt); err != nil { - return nil, err + for i := 0; ; i++ { + query, row, err = d.prepareAndQueryRow(stmt, args...) + if err == nil || i >= maxQueryRetryAttempts { + return row, err + } + + // Try to recover + if recoverErr := d.recoverFromErr(err); recoverErr != nil { + return nil, err // Unable to recover. + } } - defer p.Close() - row, err = p.QueryRow(args...), nil - return + panic("reached") } // Driver returns the underlying *sql.DB or *sql.Tx instance. func (d *database) Driver() interface{} { if tx := d.Transaction(); tx != nil { - // A transaction return tx.(*sqlTx).Tx } - return d.sess + return d.Session() } // prepareStatement converts a *exql.Statement representation into an actual // *sql.Stmt. This method will attempt to used a cached prepared statement, if // available. func (d *database) prepareStatement(stmt *exql.Statement) (*Stmt, string, error) { - if d.sess == nil && d.Transaction() == nil { + if d.Session() == nil && d.Transaction() == nil { return nil, "", db.ErrNotConnected } pc, ok := d.cachedStatements.ReadRaw(stmt) if ok { - // The statement was cached. + // This prepared statement was cached, no need to build or to prepare + // again. ps, err := pc.(*Stmt).Open() if err == nil { return ps, ps.query, nil } } - // Plain SQL query. + // Building the actual SQL query. query := d.PartialDatabase.CompileStatement(stmt) sqlStmt, err := func() (*sql.Stmt, error) { - if d.Transaction() != nil { - return d.Transaction().(*sqlTx).Prepare(query) + if tx := d.Transaction(); tx != nil { + return tx.(*sqlTx).Prepare(query) } return d.sess.Prepare(query) }() @@ -385,43 +615,15 @@ func (d *database) prepareStatement(stmt *exql.Statement) (*Stmt, string, error) var waitForConnMu sync.Mutex -// WaitForConnection tries to execute the given connectFn function, if -// connectFn returns an error, then WaitForConnection will keep trying until -// connectFn returns nil. Maximum waiting time is 5s after having acquired the -// lock. -func (d *database) WaitForConnection(connectFn func() error) error { - // This lock ensures first-come, first-served and prevents opening too many - // file descriptors. - waitForConnMu.Lock() - defer waitForConnMu.Unlock() - - // Minimum waiting time. - waitTime := time.Millisecond * 10 - - // Waitig 5 seconds for a successful connection. - for timeStart := time.Now(); time.Now().Sub(timeStart) < time.Second*5; { - err := connectFn() - if err == nil { - return nil // Connected! - } - - // Only attempt to reconnect if the error is too many clients. - if d.PartialDatabase.Err(err) == db.ErrTooManyClients { - // Sleep and try again if, and only if, the server replied with a "too - // many clients" error. - time.Sleep(waitTime) - if waitTime < time.Millisecond*500 { - // Wait a bit more next time. - waitTime = waitTime * 2 - } - continue - } - - // Return any other error immediately. +// WaitForConnection tries to execute the given connFn function, if connFn +// returns an error, then WaitForConnection will keep trying until connFn +// returns nil. Maximum waiting time is 5s after having acquired the lock. +func (d *database) WaitForConnection(connFn func() error) error { + if err := d.connect(connFn); err != nil { return err } - - return db.ErrGivingUpTryingToConnect + // Success. + return nil } // ReplaceWithDollarSign turns a SQL statament with '?' placeholders into @@ -448,16 +650,24 @@ func ReplaceWithDollarSign(in string) string { func newSessionID() uint64 { if atomic.LoadUint64(&lastSessID) == math.MaxUint64 { - atomic.StoreUint64(&lastSessID, 0) - return 0 + atomic.StoreUint64(&lastSessID, 1) + return 1 } return atomic.AddUint64(&lastSessID, 1) } func newTxID() uint64 { if atomic.LoadUint64(&lastTxID) == math.MaxUint64 { - atomic.StoreUint64(&lastTxID, 0) - return 0 + atomic.StoreUint64(&lastTxID, 1) + return 1 } return atomic.AddUint64(&lastTxID, 1) } + +func newOperationID() uint64 { + if atomic.LoadUint64(&lastOperationID) == math.MaxUint64 { + atomic.StoreUint64(&lastOperationID, 1) + return 1 + } + return atomic.AddUint64(&lastOperationID, 1) +} diff --git a/internal/sqladapter/testing/adapter.go.tpl b/internal/sqladapter/testing/adapter.go.tpl index 81774a2c..2d1fba1b 100644 --- a/internal/sqladapter/testing/adapter.go.tpl +++ b/internal/sqladapter/testing/adapter.go.tpl @@ -76,7 +76,7 @@ func TestOpenMustSucceed(t *testing.T) { assert.NoError(t, err) } -func TestPreparedStatementsCache(t *testing.T) { +func TestStressPreparedStatementsCache(t *testing.T) { sess, err := Open(settings) assert.NoError(t, err) defer sess.Close() @@ -1429,7 +1429,7 @@ func TestBuilder(t *testing.T) { assert.NotZero(t, all) } -func TestExhaustConnectionPool(t *testing.T) { +func TestExhaustConnectionPoolWithTransactions(t *testing.T) { if Adapter == "ql" { t.Skip("Currently not supported.") } @@ -1461,7 +1461,7 @@ func TestExhaustConnectionPool(t *testing.T) { // Requesting a new transaction session. start := time.Now() - tLogf("Tx: %d: NewTx") + tLogf("Tx: %d: NewTx", i) tx, err := sess.NewTx() if err != nil { tFatal(err) diff --git a/lib/sqlbuilder/fetch.go b/lib/sqlbuilder/fetch.go index b5bf577f..e3f6b1fa 100644 --- a/lib/sqlbuilder/fetch.go +++ b/lib/sqlbuilder/fetch.go @@ -81,6 +81,9 @@ func fetchRow(rows *sql.Rows, dst interface{}) error { // slice of structs given by the pointer `dst`. func fetchRows(rows *sql.Rows, dst interface{}) error { var err error + if rows == nil { + panic("rows cannot be nil") + } defer rows.Close() diff --git a/logger.go b/logger.go index 1703fb58..4f8d949f 100644 --- a/logger.go +++ b/logger.go @@ -47,8 +47,9 @@ var ( // QueryStatus represents the status of a query after being executed. type QueryStatus struct { - SessID uint64 - TxID uint64 + SessID uint64 + TxID uint64 + QueryID uint64 RowsAffected *int64 LastInsertID *int64 diff --git a/postgresql/database.go b/postgresql/database.go index d7807d99..2b89919b 100644 --- a/postgresql/database.go +++ b/postgresql/database.go @@ -153,10 +153,21 @@ func (d *database) CompileStatement(stmt *exql.Statement) string { // Err allows sqladapter to translate some known errors into generic errors. func (d *database) Err(err error) error { if err != nil { + // These errors are not exported so we have to check them by comparing + // string values. s := err.Error() - // These errors are not exported so we have to check them by they string value. - if strings.Contains(s, `too many clients`) || strings.Contains(s, `remaining connection slots are reserved`) || strings.Contains(s, `too many open`) { + switch { + case strings.Contains(s, `too many clients`), + strings.Contains(s, `remaining connection slots are reserved`), + strings.Contains(s, `too many open`): return db.ErrTooManyClients + case strings.Contains(s, `connection refused`), + strings.Contains(s, `reset by peer`), + strings.Contains(s, `is starting up`), + strings.Contains(s, `is in recovery mode`), + strings.Contains(s, `is closed`), + strings.Contains(s, `is shutting down`): + return db.ErrServerRefusedConnection } } return err @@ -194,7 +205,6 @@ func (d *database) NewLocalTransaction() (sqladapter.DatabaseTx, error) { if err := d.BaseDatabase.WaitForConnection(connFn); err != nil { return nil, err } - return sqladapter.NewTx(clone), nil } diff --git a/postgresql/local_test.go b/postgresql/local_test.go index b09e86ad..b54682bf 100644 --- a/postgresql/local_test.go +++ b/postgresql/local_test.go @@ -2,6 +2,8 @@ package postgresql import ( "database/sql" + "os" + "sync" "testing" "github.com/stretchr/testify/assert" @@ -117,3 +119,44 @@ func TestIssue210(t *testing.T) { _, err = sess.Collection("hello").Find().Count() assert.NoError(t, err) } + +// The "driver: bad connection" problem (driver.ErrBadConn) happens when the +// database process is abnormally interrupted, this can happen for a variety of +// reasons, for instance when the database process is OOM killed by the OS. +func TestDriverBadConnection(t *testing.T) { + sess := mustOpen() + defer sess.Close() + + var tMu sync.Mutex + tFatal := func(err error) { + tMu.Lock() + defer tMu.Unlock() + t.Fatal(err) + } + + const concurrentOpts = 20 + var wg sync.WaitGroup + + for i := 0; i < 100; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + _, err := sess.Collection("artist").Find().Count() + if err != nil { + tFatal(err) + } + }(i) + if i%concurrentOpts == (concurrentOpts - 1) { + // This triggers the "bad connection" problem, if you want to see that + // instead of retrying the query set maxQueryRetryAttempts to 0. + wg.Add(1) + go func() { + defer wg.Done() + sess.Query("SELECT pg_terminate_backend(pg_stat_activity.pid) FROM pg_stat_activity WHERE pid <> pg_backend_pid() AND datname = ?", os.Getenv("DB_NAME")) + }() + wg.Wait() + } + } + + wg.Wait() +}