b93e14ae88f05c6e01d25ab41d2ecb7456d769b3 / internal/db/sessions.go · 2165 bytes · raw
package db
import (
"database/sql"
"errors"
"fmt"
)
// Session is a row in the sessions table.
type Session struct {
Token string
UserID int64
ExpiresAt int64
}
// CreateSession stores a new login session for user userID.
func CreateSession(database *sql.DB, token string, userID, expiresAt int64) error {
_, err := database.Exec(
`INSERT INTO sessions (token, user_id, expires_at) VALUES (?, ?, ?)`,
token, userID, expiresAt,
)
if err != nil {
return fmt.Errorf("create session: %w", err)
}
return nil
}
// GetSession looks up a session by token, returning ErrNotFound for
// unknown or expired tokens alike.
func GetSession(database *sql.DB, token string) (Session, error) {
var session Session
err := database.QueryRow(
`SELECT token, user_id, expires_at FROM sessions
WHERE token = ? AND expires_at > unixepoch()`,
token,
).Scan(&session.Token, &session.UserID, &session.ExpiresAt)
if errors.Is(err, sql.ErrNoRows) {
return Session{}, fmt.Errorf("session: %w", ErrNotFound)
}
if err != nil {
return Session{}, fmt.Errorf("get session: %w", err)
}
return session, nil
}
// DeleteSession removes a session at logout.
func DeleteSession(database *sql.DB, token string) error {
_, err := database.Exec(`DELETE FROM sessions WHERE token = ?`, token)
if err != nil {
return fmt.Errorf("delete session: %w", err)
}
return nil
}
// DeleteOtherSessions revokes every session for a user except keepToken,
// so a password change signs other browsers out.
func DeleteOtherSessions(database *sql.DB, userID int64, keepToken string) error {
_, err := database.Exec(
`DELETE FROM sessions WHERE user_id = ? AND token != ?`, userID, keepToken)
if err != nil {
return fmt.Errorf("delete other sessions for user %d: %w", userID, err)
}
return nil
}
// DeleteExpiredSessions purges sessions whose expiry has passed; callers
// use it as opportunistic housekeeping, so no error on zero rows.
func DeleteExpiredSessions(database *sql.DB) error {
if _, err := database.Exec(`DELETE FROM sessions WHERE expires_at <= unixepoch()`); err != nil {
return fmt.Errorf("delete expired sessions: %w", err)
}
return nil
}