增加了一些功能

This commit is contained in:
2026-06-19 23:08:26 +08:00
parent 6dd576158c
commit aedc199777
9 changed files with 502 additions and 10 deletions
+143
View File
@@ -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
}