feat: 实现 Admin 首次初始化与网页登录 (#50)

This commit is contained in:
chengma
2026-08-09 13:36:41 +08:00
parent 1064ffed59
commit 0a7db404a1
17 changed files with 1140 additions and 12 deletions
+153
View File
@@ -0,0 +1,153 @@
// Admin 网页认证:首次管理员、密码校验、随机 Session 和过期判断。
package service
import (
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"
"golang.org/x/crypto/bcrypt"
"cmautobuy/admin/model"
"cmautobuy/admin/repository"
)
const (
WebSessionDuration = 12 * time.Hour
minimumPasswordLen = 8
maxUsernameLen = 64
)
var (
ErrInvalidCredentials = errors.New("用户名或密码错误")
ErrUnauthenticated = errors.New("未登录或登录已过期")
// 用户不存在时也跑一次 bcrypt,避免响应时间直接泄露“这个用户名存在”。
dummyPasswordHash, _ = bcrypt.GenerateFromPassword([]byte("not-a-real-password"), bcrypt.DefaultCost)
)
// HasUsers 判断首次初始化入口是否已经永久关闭。
func HasUsers(db *sql.DB) (bool, error) {
count, err := repository.CountUsers(db)
return count > 0, err
}
// SetupInitialAdmin 校验表单、哈希密码并在事务中创建首位管理员。
func SetupInitialAdmin(db *sql.DB, username, password, confirmation string, now time.Time) error {
username = strings.TrimSpace(username)
if username == "" {
return fmt.Errorf("用户名不能为空")
}
if len([]rune(username)) > maxUsernameLen {
return fmt.Errorf("用户名不能超过 %d 个字符", maxUsernameLen)
}
if len([]rune(password)) < minimumPasswordLen {
return fmt.Errorf("密码至少需要 %d 个字符", minimumPasswordLen)
}
// bcrypt 最多接受 72 字节。中文等字符可能占多个字节,因此不能只靠
// HTML 的 maxlength;服务端需要在哈希前给出可理解的校验错误。
if len([]byte(password)) > 72 {
return fmt.Errorf("密码不能超过 72 个字节")
}
if password != confirmation {
return fmt.Errorf("两次输入的密码不一致")
}
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return fmt.Errorf("生成密码哈希失败: %w", err)
}
userID, err := randomID("USR-", 16)
if err != nil {
return err
}
at := now.UTC().Format(model.TimeLayout)
return repository.CreateInitialAdmin(db, model.User{
UserID: userID, Username: username, PasswordHash: string(hash),
Role: model.RoleAdmin, Status: model.UserActive,
PasswordChangedAt: at, CreatedAt: at, UpdatedAt: at,
})
}
// Login 校验统一凭据并创建一个固定 12 小时有效的 Session。
// 返回的 token 原文只交给 Cookie,数据库仅保存 SHA-256。
func Login(db *sql.DB, username, password string, now time.Time) (token string, user *model.User, expiresAt time.Time, err error) {
username = strings.TrimSpace(username)
user, findErr := repository.FindUserByUsername(db, username)
hash := dummyPasswordHash
if findErr == nil {
hash = []byte(user.PasswordHash)
} else if !errors.Is(findErr, repository.ErrUserNotFound) {
return "", nil, time.Time{}, findErr
}
passwordOK := bcrypt.CompareHashAndPassword(hash, []byte(password)) == nil
if findErr != nil || !passwordOK || user.Status != model.UserActive {
return "", nil, time.Time{}, ErrInvalidCredentials
}
tokenBytes := make([]byte, 32)
if _, err := rand.Read(tokenBytes); err != nil {
return "", nil, time.Time{}, fmt.Errorf("生成 Web Session 失败: %w", err)
}
token = base64.RawURLEncoding.EncodeToString(tokenBytes)
expiresAt = now.Add(WebSessionDuration).UTC()
at := now.UTC().Format(model.TimeLayout)
if err := repository.DeleteExpiredSessions(db, at); err != nil {
return "", nil, time.Time{}, err
}
if err := repository.CreateLoginSession(db, model.WebSession{
SessionHash: SessionTokenHash(token), UserID: user.UserID,
ExpiresAt: expiresAt.Format(model.TimeLayout), CreatedAt: at, LastSeenAt: at,
}, at); err != nil {
return "", nil, time.Time{}, err
}
return token, user, expiresAt, nil
}
// Authenticate 验证 Cookie Token、固定过期时间和账号状态。
func Authenticate(db *sql.DB, token string, now time.Time) (*model.User, error) {
if strings.TrimSpace(token) == "" {
return nil, ErrUnauthenticated
}
session, user, err := repository.FindUserBySessionHash(db, SessionTokenHash(token))
if errors.Is(err, repository.ErrSessionNotFound) {
return nil, ErrUnauthenticated
}
if err != nil {
return nil, err
}
expiresAt, ok := model.ParseISO(session.ExpiresAt)
if !ok || !now.UTC().Before(expiresAt) || user.Status != model.UserActive {
_ = repository.DeleteSession(db, session.SessionHash)
return nil, ErrUnauthenticated
}
return user, nil
}
// Logout 撤销当前 Token 对应的服务端 Session。空 Token 也视为成功。
func Logout(db *sql.DB, token string) error {
if token == "" {
return nil
}
return repository.DeleteSession(db, SessionTokenHash(token))
}
// SessionTokenHash 返回数据库可保存的 Token SHA-256 十六进制文本。
func SessionTokenHash(token string) string {
sum := sha256.Sum256([]byte(token))
return hex.EncodeToString(sum[:])
}
func randomID(prefix string, byteCount int) (string, error) {
buf := make([]byte, byteCount)
if _, err := rand.Read(buf); err != nil {
return "", fmt.Errorf("生成用户编号失败: %w", err)
}
return prefix + hex.EncodeToString(buf), nil
}
+164
View File
@@ -0,0 +1,164 @@
package service
import (
"errors"
"strings"
"sync"
"testing"
"time"
"golang.org/x/crypto/bcrypt"
"cmautobuy/admin/repository"
)
func TestSetupInitialAdmin_保存哈希且入口永久关闭(t *testing.T) {
db := newSyncTestDB(t)
now := time.Date(2026, 8, 9, 6, 0, 0, 0, time.UTC)
if err := SetupInitialAdmin(db, "admin", "safe-password", "safe-password", now); err != nil {
t.Fatalf("初始化管理员失败: %v", err)
}
user, err := repository.FindUserByUsername(db, "ADMIN") // username 是 NOCASE
if err != nil {
t.Fatalf("读取管理员失败: %v", err)
}
if user.PasswordHash == "safe-password" || strings.Contains(user.PasswordHash, "safe-password") {
t.Fatal("数据库不能保存或包含明文密码")
}
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte("safe-password")); err != nil {
t.Fatalf("密码哈希不能验证原密码: %v", err)
}
if err := SetupInitialAdmin(db, "second", "other-password", "other-password", now); !errors.Is(err, repository.ErrUsersAlreadyExist) {
t.Fatalf("已有用户后初始化入口应永久关闭,实际 %v", err)
}
}
func TestSetupInitialAdmin_校验密码字符数和Bcrypt字节上限(t *testing.T) {
for _, test := range []struct {
name string
password string
want string
}{
{"少于八个字符", "1234567", "至少需要 8 个字符"},
{"超过bcrypt字节上限", strings.Repeat("密", 25), "不能超过 72 个字节"},
} {
t.Run(test.name, func(t *testing.T) {
db := newSyncTestDB(t)
err := SetupInitialAdmin(db, "admin", test.password, test.password, time.Now())
if err == nil || !strings.Contains(err.Error(), test.want) {
t.Fatalf("校验错误 = %v,期望包含 %q", err, test.want)
}
})
}
}
func TestSetupInitialAdmin_并发最多一个成功(t *testing.T) {
db := newSyncTestDB(t)
start := make(chan struct{})
errorsOut := make(chan error, 2)
var wg sync.WaitGroup
for _, username := range []string{"admin-a", "admin-b"} {
wg.Add(1)
go func(name string) {
defer wg.Done()
<-start
errorsOut <- SetupInitialAdmin(db, name, "safe-password", "safe-password", time.Now())
}(username)
}
close(start)
wg.Wait()
close(errorsOut)
successes := 0
for err := range errorsOut {
if err == nil {
successes++
continue
}
if !errors.Is(err, repository.ErrUsersAlreadyExist) {
t.Errorf("失败请求应明确提示已初始化,实际 %v", err)
}
}
if successes != 1 {
t.Fatalf("并发初始化应恰好一个成功,实际 %d", successes)
}
count, _ := repository.CountUsers(db)
if count != 1 {
t.Fatalf("并发初始化后数据库应只有一个用户,实际 %d", count)
}
}
func TestLoginAuthenticateLogout_数据库只存Token哈希(t *testing.T) {
db := newSyncTestDB(t)
now := time.Date(2026, 8, 9, 6, 0, 0, 0, time.UTC)
if err := SetupInitialAdmin(db, "admin", "safe-password", "safe-password", now); err != nil {
t.Fatal(err)
}
token, user, expiresAt, err := Login(db, "admin", "safe-password", now)
if err != nil {
t.Fatalf("登录失败: %v", err)
}
if len(token) < 40 || expiresAt.Sub(now) != WebSessionDuration {
t.Fatalf("Session Token 或 12 小时过期时间不正确")
}
var stored string
if err := db.QueryRow(`SELECT session_hash FROM web_sessions`).Scan(&stored); err != nil {
t.Fatal(err)
}
if stored == token || stored != SessionTokenHash(token) || len(stored) != 64 {
t.Fatalf("数据库应只保存 64 位 SHA-256,stored=%q token=%q", stored, token)
}
got, err := Authenticate(db, token, now.Add(time.Hour))
if err != nil || got.UserID != user.UserID {
t.Fatalf("有效 Session 认证失败: user=%+v err=%v", got, err)
}
if err := Logout(db, token); err != nil {
t.Fatalf("退出失败: %v", err)
}
if _, err := Authenticate(db, token, now.Add(time.Hour)); !errors.Is(err, ErrUnauthenticated) {
t.Fatalf("退出后旧 Token 不应继续有效,实际 %v", err)
}
}
func TestLogin_错误凭据和禁用账号统一提示(t *testing.T) {
db := newSyncTestDB(t)
if err := SetupInitialAdmin(db, "admin", "safe-password", "safe-password", time.Now()); err != nil {
t.Fatal(err)
}
for name, password := range map[string]string{
"missing": "safe-password",
"admin": "wrong-password",
} {
if _, _, _, err := Login(db, name, password, time.Now()); !errors.Is(err, ErrInvalidCredentials) {
t.Errorf("%s 登录错误应统一,实际 %v", name, err)
}
}
if _, err := db.Exec(`UPDATE users SET status = 'disabled' WHERE username = 'admin'`); err != nil {
t.Fatal(err)
}
if _, _, _, err := Login(db, "admin", "safe-password", time.Now()); !errors.Is(err, ErrInvalidCredentials) {
t.Fatalf("禁用账号也应使用统一提示,实际 %v", err)
}
}
func TestAuthenticate_过期Session被删除(t *testing.T) {
db := newSyncTestDB(t)
now := time.Date(2026, 8, 9, 6, 0, 0, 0, time.UTC)
if err := SetupInitialAdmin(db, "admin", "safe-password", "safe-password", now); err != nil {
t.Fatal(err)
}
token, _, _, err := Login(db, "admin", "safe-password", now)
if err != nil {
t.Fatal(err)
}
if _, err := Authenticate(db, token, now.Add(WebSessionDuration)); !errors.Is(err, ErrUnauthenticated) {
t.Fatalf("到期时刻 Session 必须失效,实际 %v", err)
}
var count int
if err := db.QueryRow(`SELECT COUNT(*) FROM web_sessions`).Scan(&count); err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("过期 Session 应被清理,实际剩 %d", count)
}
}