package users import ( "crypto/rand" "database/sql" "encoding/hex" "errors" "fmt" "strings" "time" "golang.org/x/crypto/bcrypt" _ "modernc.org/sqlite" ) var ErrNotFound = errors.New("user not found") var ErrUserExists = errors.New("user already exists") type User struct { ID int64 Username string Role Role } type Session struct { Token string UserID int64 Username string Role Role ExpiresAt time.Time } type Store struct { db *sql.DB } func Open(path string) (*Store, error) { db, err := sql.Open("sqlite", path) if err != nil { return nil, err } db.SetMaxOpenConns(1) // SQLite doesn't support concurrent writers s := &Store{db: db} if err := s.migrate(); err != nil { db.Close() return nil, err } return s, nil } func (s *Store) Close() error { return s.db.Close() } // migrations is an ordered list of schema changes. Each entry is a slice of // SQL statements to execute in a transaction. Add new entries at the end only. var migrations = [][]string{ // v1: initial schema { `CREATE TABLE IF NOT EXISTS users ( id INTEGER PRIMARY KEY AUTOINCREMENT, username TEXT NOT NULL UNIQUE, password_hash TEXT NOT NULL, can_upload INTEGER NOT NULL DEFAULT 0 )`, `CREATE TABLE IF NOT EXISTS sessions ( token TEXT PRIMARY KEY, user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, expires_at INTEGER NOT NULL )`, }, // v2: replace can_upload with role column; migrate existing data { `ALTER TABLE users ADD COLUMN role TEXT NOT NULL DEFAULT 'reader'`, `UPDATE users SET role = CASE WHEN can_upload = 1 THEN 'uploader' ELSE 'reader' END`, }, } func (s *Store) migrate() error { if _, err := s.db.Exec( `CREATE TABLE IF NOT EXISTS _schema_version (version INTEGER NOT NULL)`, ); err != nil { return err } var version int if err := s.db.QueryRow( `SELECT COALESCE(MAX(version), 0) FROM _schema_version`, ).Scan(&version); err != nil { return err } for i, stmts := range migrations { v := i + 1 if v <= version { continue } tx, err := s.db.Begin() if err != nil { return err } for _, stmt := range stmts { if _, err := tx.Exec(stmt); err != nil { _ = tx.Rollback() return fmt.Errorf("migration %d: %w", v, err) } } if _, err := tx.Exec(`INSERT INTO _schema_version (version) VALUES (?)`, v); err != nil { _ = tx.Rollback() return err } if err := tx.Commit(); err != nil { return err } } return nil } // CreateUser adds a new user with the given role. Returns ErrUserExists if the // username is already taken. func (s *Store) CreateUser(username, password string, role Role) (*User, error) { hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost) if err != nil { return nil, err } res, err := s.db.Exec( `INSERT INTO users (username, password_hash, role) VALUES (?, ?, ?)`, username, string(hash), role, ) if err != nil { if isUnique(err) { return nil, ErrUserExists } return nil, err } id, _ := res.LastInsertId() return &User{ID: id, Username: username, Role: role}, nil } // SetRole changes the role of an existing user. func (s *Store) SetRole(username string, role Role) error { res, err := s.db.Exec( `UPDATE users SET role = ? WHERE username = ?`, role, username, ) if err != nil { return err } if n, _ := res.RowsAffected(); n == 0 { return ErrNotFound } return nil } // DeleteUser removes a user by username. func (s *Store) DeleteUser(username string) error { res, err := s.db.Exec(`DELETE FROM users WHERE username = ?`, username) if err != nil { return err } if n, _ := res.RowsAffected(); n == 0 { return ErrNotFound } return nil } // ListUsers returns all users ordered by username. func (s *Store) ListUsers() ([]User, error) { rows, err := s.db.Query(`SELECT id, username, role FROM users ORDER BY username`) if err != nil { return nil, err } defer rows.Close() var users []User for rows.Next() { var u User var roleStr string if err := rows.Scan(&u.ID, &u.Username, &roleStr); err != nil { return nil, err } u.Role = Role(roleStr) users = append(users, u) } return users, rows.Err() } // Authenticate verifies credentials and returns the matching user, or nil on // wrong username/password. func (s *Store) Authenticate(username, password string) (*User, error) { var u User var hash, roleStr string err := s.db.QueryRow( `SELECT id, username, password_hash, role FROM users WHERE username = ?`, username, ).Scan(&u.ID, &u.Username, &hash, &roleStr) if errors.Is(err, sql.ErrNoRows) { return nil, nil } if err != nil { return nil, err } if bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) != nil { return nil, nil } u.Role = Role(roleStr) return &u, nil } // CreateSession issues a new session token for a user (TTL: 30 days). func (s *Store) CreateSession(u *User) (*Session, error) { token, err := randomToken() if err != nil { return nil, err } exp := time.Now().Add(30 * 24 * time.Hour) if _, err = s.db.Exec( `INSERT INTO sessions (token, user_id, expires_at) VALUES (?, ?, ?)`, token, u.ID, exp.Unix(), ); err != nil { return nil, err } return &Session{ Token: token, UserID: u.ID, Username: u.Username, Role: u.Role, ExpiresAt: exp, }, nil } // LookupSession returns the session if valid, nil if not found or expired. func (s *Store) LookupSession(token string) (*Session, error) { var sess Session var roleStr string var expUnix int64 err := s.db.QueryRow(` SELECT s.token, s.user_id, u.username, u.role, s.expires_at FROM sessions s JOIN users u ON u.id = s.user_id WHERE s.token = ? `, token).Scan(&sess.Token, &sess.UserID, &sess.Username, &roleStr, &expUnix) if errors.Is(err, sql.ErrNoRows) { return nil, nil } if err != nil { return nil, err } sess.ExpiresAt = time.Unix(expUnix, 0) if time.Now().After(sess.ExpiresAt) { _, _ = s.db.Exec(`DELETE FROM sessions WHERE token = ?`, token) return nil, nil } sess.Role = Role(roleStr) return &sess, nil } // DeleteSession invalidates a session token (logout). func (s *Store) DeleteSession(token string) error { _, err := s.db.Exec(`DELETE FROM sessions WHERE token = ?`, token) return err } func randomToken() (string, error) { b := make([]byte, 32) if _, err := rand.Read(b); err != nil { return "", fmt.Errorf("random token: %w", err) } return hex.EncodeToString(b), nil } func isUnique(err error) bool { return err != nil && strings.Contains(strings.ToLower(err.Error()), "unique") }