144 lines
3.7 KiB
Go
144 lines
3.7 KiB
Go
package db
|
|
|
|
import (
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"database/sql"
|
|
"encoding/hex"
|
|
"fmt"
|
|
"os"
|
|
"path/filepath"
|
|
"time"
|
|
|
|
_ "modernc.org/sqlite"
|
|
)
|
|
|
|
var DB *sql.DB
|
|
|
|
func Init() error {
|
|
dir, _ := os.Executable()
|
|
dbPath := filepath.Join(filepath.Dir(dir), "data", "game.db")
|
|
os.MkdirAll(filepath.Dir(dbPath), 0755)
|
|
|
|
var err error
|
|
DB, err = sql.Open("sqlite", dbPath)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
_, err = DB.Exec(`
|
|
CREATE TABLE IF NOT EXISTS users (
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
username TEXT UNIQUE NOT NULL,
|
|
password_hash TEXT NOT NULL,
|
|
salt TEXT NOT NULL,
|
|
nickname TEXT NOT NULL,
|
|
bio TEXT DEFAULT '',
|
|
avatar TEXT DEFAULT '',
|
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
|
last_login TIMESTAMP
|
|
);
|
|
CREATE TABLE IF NOT EXISTS sessions (
|
|
token TEXT PRIMARY KEY,
|
|
user_id INTEGER NOT NULL,
|
|
username TEXT NOT NULL,
|
|
nickname TEXT NOT NULL,
|
|
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
|
FOREIGN KEY (user_id) REFERENCES users(id)
|
|
);
|
|
`)
|
|
return err
|
|
}
|
|
|
|
type User struct {
|
|
ID int64
|
|
Username string
|
|
Nickname string
|
|
Bio string
|
|
Avatar string
|
|
}
|
|
|
|
func hashPassword(password, salt string) string {
|
|
h := sha256.Sum256([]byte(password + salt))
|
|
return hex.EncodeToString(h[:])
|
|
}
|
|
|
|
func generateSalt() string {
|
|
b := make([]byte, 16)
|
|
rand.Read(b)
|
|
return hex.EncodeToString(b)
|
|
}
|
|
|
|
func generateToken() string {
|
|
b := make([]byte, 32)
|
|
rand.Read(b)
|
|
return hex.EncodeToString(b)
|
|
}
|
|
|
|
func Register(username, password, nickname string) (*User, error) {
|
|
if len(username) < 2 || len(password) < 4 {
|
|
return nil, fmt.Errorf("用户名至少2位,密码至少4位")
|
|
}
|
|
if nickname == "" {
|
|
nickname = username
|
|
}
|
|
salt := generateSalt()
|
|
hash := hashPassword(password, salt)
|
|
res, err := DB.Exec("INSERT INTO users (username, password_hash, salt, nickname) VALUES (?, ?, ?, ?)",
|
|
username, hash, salt, nickname)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("用户名已存在")
|
|
}
|
|
id, _ := res.LastInsertId()
|
|
return &User{ID: id, Username: username, Nickname: nickname}, nil
|
|
}
|
|
|
|
func Login(username, password string) (string, *User, error) {
|
|
var id int64
|
|
var hash, salt, nickname, bio, avatar string
|
|
err := DB.QueryRow("SELECT id, password_hash, salt, nickname, bio, avatar FROM users WHERE username = ?",
|
|
username).Scan(&id, &hash, &salt, &nickname, &bio, &avatar)
|
|
if err != nil {
|
|
return "", nil, fmt.Errorf("用户不存在")
|
|
}
|
|
if hashPassword(password, salt) != hash {
|
|
return "", nil, fmt.Errorf("密码错误")
|
|
}
|
|
DB.Exec("UPDATE users SET last_login = ? WHERE id = ?", time.Now(), id)
|
|
|
|
token := generateToken()
|
|
DB.Exec("INSERT OR REPLACE INTO sessions (token, user_id, username, nickname) VALUES (?, ?, ?, ?)",
|
|
token, id, username, nickname)
|
|
|
|
return token, &User{ID: id, Username: username, Nickname: nickname, Bio: bio, Avatar: avatar}, nil
|
|
}
|
|
|
|
func ValidateToken(token string) (*User, error) {
|
|
var id int64
|
|
var username, nickname string
|
|
err := DB.QueryRow("SELECT user_id, username, nickname FROM sessions WHERE token = ?", token).
|
|
Scan(&id, &username, &nickname)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("无效的令牌")
|
|
}
|
|
return &User{ID: id, Username: username, Nickname: nickname}, nil
|
|
}
|
|
|
|
func Logout(token string) {
|
|
DB.Exec("DELETE FROM sessions WHERE token = ?", token)
|
|
}
|
|
|
|
func UpdateProfile(userID int64, nickname, bio, avatar string) error {
|
|
_, err := DB.Exec("UPDATE users SET nickname=?, bio=?, avatar=? WHERE id=?", nickname, bio, avatar, userID)
|
|
return err
|
|
}
|
|
|
|
func GetProfile(username string) (*User, error) {
|
|
var u User
|
|
err := DB.QueryRow("SELECT id, username, nickname, bio, avatar FROM users WHERE username=?", username).
|
|
Scan(&u.ID, &u.Username, &u.Nickname, &u.Bio, &u.Avatar)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("用户不存在")
|
|
}
|
|
return &u, nil
|
|
}
|