josie / simplegit

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
}

// CountExpiredSessions returns the number of sessions past their expiry.
func CountExpiredSessions(database *sql.DB) (int, error) {
	var n int
	if err := database.QueryRow(`SELECT count(*) FROM sessions WHERE expires_at <= unixepoch()`).Scan(&n); err != nil {
		return 0, fmt.Errorf("count expired sessions: %w", err)
	}
	return n, 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
}