gowiki/userdb.go

278 lines
5.8 KiB
Go

package main
import (
"bytes"
"encoding/gob"
"errors"
"github.com/boltdb/bolt"
"golang.org/x/crypto/bcrypt"
"golang.org/x/crypto/sha3"
)
type User struct {
Username string
Password string
Admin bool
Locked int
}
func validatePassword(pw string) (b []byte) {
b = []byte(pw)
if len(pw) <= 56 {
return
}
hash := sha3.New384()
hash.Write(b)
b = hash.Sum(nil)
return
}
const (
userTable = "user"
bcryptCost = 13
)
var (
ErrNotFound = errors.New("db: key not found")
ErrExists = errors.New("db: key already exists")
)
type BoltStore struct {
Marshaler
db *bolt.DB
}
func NewBoltStore(file string) (*BoltStore, error) {
db, err := bolt.Open(file, 0666, nil)
if err != nil {
return nil, err
}
if err = db.Update(func(tx *bolt.Tx) error {
if _, err := tx.CreateBucketIfNotExists([]byte([]byte(userTable))); err != nil {
return err
}
return nil
}); err != nil {
return nil, err
}
bs := &BoltStore{NewGOB(), db}
bs.UnlockUser(defUser)
return bs, nil
}
func (bs *BoltStore) Close() error {
return bs.db.Close()
}
func (bs *BoltStore) GetUsers() (us []User, err error) {
err = bs.db.View(func(tx *bolt.Tx) error {
return tx.Bucket([]byte(userTable)).ForEach(func(k, v []byte) error {
u := new(User)
if err := bs.Unmarshal(v, u); err != nil {
return err
}
u.Password = ""
us = append(us, *u)
return nil
})
})
return
}
func (bs *BoltStore) GetUserWithoutPassword(username string) (u User, err error) {
err = bs.db.View(func(tx *bolt.Tx) error {
v := tx.Bucket([]byte(userTable)).Get([]byte(username))
if v == nil {
return ErrNotFound
}
return bs.Unmarshal(v, &u)
})
u.Password = ""
return
}
func (bs *BoltStore) GetUser(username string) (u User, err error) {
err = bs.db.View(func(tx *bolt.Tx) error {
v := tx.Bucket([]byte(userTable)).Get([]byte(username))
if v == nil {
return ErrNotFound
}
return bs.Unmarshal(v, &u)
})
return
}
func (bs *BoltStore) UserIsAdmin(username string) (admin bool) {
bs.db.View(func(tx *bolt.Tx) error {
u, err := bs.GetUser(username)
if err != nil {
return err
}
admin = u.Admin
return nil
})
return
}
func (bs *BoltStore) UserExists(username string) (exists bool) {
bs.db.View(func(tx *bolt.Tx) error {
exists = tx.Bucket([]byte(userTable)).Get([]byte(username)) != nil
return nil
})
return
}
func (bs *BoltStore) CreateUser(username, password string, admin bool) error {
return bs.db.Update(func(tx *bolt.Tx) error {
if bs.UserExists(username) {
return ErrExists
}
hash, err := bcrypt.GenerateFromPassword(validatePassword(password), bcryptCost)
v, err := bs.Marshal(&User{
Username: username,
Password: string(hash),
Admin: admin,
})
if err != nil {
return err
}
return tx.Bucket([]byte(userTable)).Put([]byte(username), v)
})
}
func (bs *BoltStore) Login(username, password string) error {
var e error
err := bs.db.Update(func(tx *bolt.Tx) error {
us, err := bs.GetUser(username)
if err != nil {
return err
}
if us.Locked == 3 {
return errors.New("db: user locked")
}
b := tx.Bucket([]byte(userTable))
if err := bcrypt.CompareHashAndPassword([]byte(us.Password), validatePassword(password)); err != nil {
us.Locked++
v, err := bs.Marshal(us)
if err != nil {
return err
}
if err = b.Put([]byte(username), v); err != nil {
return err
}
e = errors.New("db: wrong login")
return nil
}
us.Locked = 0
v, err := bs.Marshal(us)
if err != nil {
return err
}
return b.Put([]byte(username), v)
})
if e != nil {
return e
}
return err
}
func (bs *BoltStore) UnlockUser(username string) error {
return bs.db.Update(func(tx *bolt.Tx) error {
u, err := bs.GetUser(username)
if err != nil {
return err
}
if u.Locked < 3 {
return nil
}
u.Locked = 0
v, err := bs.Marshal(u)
if err != nil {
return err
}
return tx.Bucket([]byte(userTable)).Put([]byte(username), v)
})
}
func (bs *BoltStore) UpdateUserPassword(username, password string) error {
return bs.db.Update(func(tx *bolt.Tx) error {
if password != "" {
us, err := bs.GetUser(username)
if err != nil {
return err
}
hash, err := bcrypt.GenerateFromPassword(validatePassword(password), bcryptCost)
if err != nil {
return err
}
us.Password = string(hash)
v, err := bs.Marshal(us)
if err != nil {
return err
}
return tx.Bucket([]byte(userTable)).Put([]byte(username), v)
}
return nil
})
}
func (bs *BoltStore) AdminUpdateUser(username, password string, admin bool) error {
if username == defUser && !admin {
return errors.New("user '" + defUser + "' cannot lose admin privileges")
}
return bs.db.Update(func(tx *bolt.Tx) error {
us, err := bs.GetUser(username)
if err != nil {
return err
}
us.Admin = admin
if password != "" {
hash, err := bcrypt.GenerateFromPassword(validatePassword(password), bcryptCost)
if err != nil {
return err
}
us.Password = string(hash)
}
v, err := bs.Marshal(us)
if err != nil {
return err
}
return tx.Bucket([]byte(userTable)).Put([]byte(username), v)
})
}
func (bs *BoltStore) DeleteUser(username string) error {
if username == defUser {
return errors.New("user '" + defUser + "' cannot be deleted")
}
return bs.db.Update(func(tx *bolt.Tx) error {
if bs.UserExists(username) {
return tx.Bucket([]byte(userTable)).Delete([]byte(username))
}
return ErrNotFound
})
}
type Marshaler interface {
Marshal(v interface{}) ([]byte, error)
Unmarshal(data []byte, v interface{}) error
}
type gobMarshaler struct{}
func NewGOB() Marshaler {
return gobMarshaler{}
}
func (gobMarshaler) Marshal(v interface{}) ([]byte, error) {
b := new(bytes.Buffer)
err := gob.NewEncoder(b).Encode(v)
if err != nil {
return nil, err
}
return b.Bytes(), nil
}
func (gobMarshaler) Unmarshal(data []byte, v interface{}) error {
return gob.NewDecoder(bytes.NewBuffer(data)).Decode(v)
}