增加了一些功能
This commit is contained in:
+143
@@ -0,0 +1,143 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user