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) )