Files
plumber/internal/store/sessions.go
T

155 lines
3.5 KiB
Go

package store
import (
"database/sql"
"log"
"time"
"github.com/alexedwards/scs/postgresstore"
"github.com/alexedwards/scs/v2"
)
const sessionsSchemaSQLite = `
CREATE TABLE IF NOT EXISTS sessions (
token TEXT PRIMARY KEY,
data BLOB NOT NULL,
expiry REAL NOT NULL
);
CREATE INDEX IF NOT EXISTS sessions_expiry_idx ON sessions(expiry);
`
const sessionsSchemaPostgres = `
CREATE TABLE IF NOT EXISTS sessions (
token TEXT PRIMARY KEY,
data BYTEA NOT NULL,
expiry TIMESTAMPTZ NOT NULL
);
CREATE INDEX IF NOT EXISTS sessions_expiry_idx ON sessions (expiry);
`
func applySessionsSchema(db *sql.DB, dialect string) error {
schema := sessionsSchemaSQLite
if dialect == dialectPostgres {
schema = sessionsSchemaPostgres
}
return applySchema(db, schema)
}
type sessionStopper interface {
StopCleanup()
}
// SessionStore returns the scs store backed by this database.
func (s *Store) SessionStore() scs.Store {
return s.sessionStore
}
func (s *Store) initSessionStore(cleanupInterval time.Duration) {
switch s.dialect {
case dialectPostgres:
ps := postgresstore.NewWithCleanupInterval(s.db, cleanupInterval)
s.sessionStore = ps
s.sessionStopper = ps
default:
ss := newSQLiteSessionStore(s.db, cleanupInterval)
s.sessionStore = ss
s.sessionStopper = ss
}
}
// sqliteSessionStore is a modernc-safe scs.Store (uses ? placeholders).
type sqliteSessionStore struct {
db *sql.DB
stopCleanup chan bool
}
func newSQLiteSessionStore(db *sql.DB, cleanupInterval time.Duration) *sqliteSessionStore {
s := &sqliteSessionStore{db: db}
if cleanupInterval > 0 {
s.stopCleanup = make(chan bool)
go s.startCleanup(cleanupInterval)
}
return s
}
func (s *sqliteSessionStore) Find(token string) ([]byte, bool, error) {
var b []byte
err := s.db.QueryRow(
`SELECT data FROM sessions WHERE token = ? AND julianday('now') < expiry`,
token,
).Scan(&b)
if err == sql.ErrNoRows {
return nil, false, nil
}
if err != nil {
return nil, false, err
}
return b, true, nil
}
func (s *sqliteSessionStore) Commit(token string, b []byte, expiry time.Time) error {
_, err := s.db.Exec(
`REPLACE INTO sessions (token, data, expiry) VALUES (?, ?, julianday(?))`,
token,
b,
expiry.UTC().Format("2006-01-02T15:04:05.999"),
)
return err
}
func (s *sqliteSessionStore) Delete(token string) error {
_, err := s.db.Exec(`DELETE FROM sessions WHERE token = ?`, token)
return err
}
func (s *sqliteSessionStore) All() (map[string][]byte, error) {
rows, err := s.db.Query(`SELECT token, data FROM sessions WHERE julianday('now') < expiry`)
if err != nil {
return nil, err
}
defer rows.Close()
out := make(map[string][]byte)
for rows.Next() {
var token string
var data []byte
if err := rows.Scan(&token, &data); err != nil {
return nil, err
}
out[token] = data
}
return out, rows.Err()
}
func (s *sqliteSessionStore) startCleanup(interval time.Duration) {
ticker := time.NewTicker(interval)
for {
select {
case <-ticker.C:
if err := s.deleteExpired(); err != nil {
log.Println(err)
}
case <-s.stopCleanup:
ticker.Stop()
return
}
}
}
func (s *sqliteSessionStore) StopCleanup() {
if s.stopCleanup != nil {
s.stopCleanup <- true
}
}
func (s *sqliteSessionStore) deleteExpired() error {
_, err := s.db.Exec(`DELETE FROM sessions WHERE expiry < julianday('now')`)
return err
}
// Ensure interface compliance.
var (
_ scs.Store = (*sqliteSessionStore)(nil)
_ scs.IterableStore = (*sqliteSessionStore)(nil)
_ sessionStopper = (*sqliteSessionStore)(nil)
)