155 lines
3.5 KiB
Go
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)
|
|
)
|