package database import ( "context" "database/sql" sqldriver "database/sql/driver" "errors" "net/http" ) var dbCtxKey = &contextKey{"database"} type contextKey struct { name string } func Middleware(db *sql.DB) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ctx := Context(r.Context(), db) r = r.WithContext(ctx) next.ServeHTTP(w, r) }) } } func Context(ctx context.Context, db *sql.DB) context.Context { return context.WithValue(ctx, dbCtxKey, db) } func ForContext(ctx context.Context) (*sql.Conn, error) { raw, ok := ctx.Value(dbCtxKey).(*sql.DB) if !ok { panic(errors.New("Invalid database context")) } return raw.Conn(ctx) } func DBForContext(ctx context.Context) *sql.DB { raw, ok := ctx.Value(dbCtxKey).(*sql.DB) if !ok { panic(errors.New("Invalid database context")) } return raw } func WithTx(ctx context.Context, opts *sql.TxOptions, fn func(tx *sql.Tx) error) error { db := DBForContext(ctx) tx, err := db.BeginTx(ctx, opts) if err != nil { return err } defer tx.Rollback() err = fn(tx) var txErr error if err != nil { txErr = tx.Rollback() } else { txErr = tx.Commit() } if errors.Is(err, context.Canceled) && errors.Is(txErr, sqldriver.ErrBadConn) { // When a query fails because the context has been canceled, pq will // return "driver: bad connection" from tx.Rollback. Do not panic in // this case. See https://github.com/lib/pq/issues/1137 return err } if txErr != nil && txErr != sql.ErrTxDone { panic(txErr) } return err }