215 lines
7.1 KiB
Go
215 lines
7.1 KiB
Go
// 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 = 6
|
|
maxUsernameLen = 64
|
|
)
|
|
|
|
var (
|
|
ErrInvalidCredentials = errors.New("用户名或密码错误")
|
|
ErrUnauthenticated = errors.New("未登录或登录已过期")
|
|
ErrForbidden = errors.New("没有权限执行此操作")
|
|
|
|
// 用户不存在时也跑一次 bcrypt,避免响应时间直接泄露“这个用户名存在”。
|
|
dummyPasswordHash, _ = bcrypt.GenerateFromPassword([]byte("not-a-real-password"), bcrypt.DefaultCost)
|
|
)
|
|
|
|
type validationError struct {
|
|
field string
|
|
message string
|
|
}
|
|
|
|
func (e *validationError) Error() string { return e.message }
|
|
|
|
// IsValidationError 让 Handler 区分可直接展示的表单错误和不可泄露的内部错误。
|
|
func IsValidationError(err error) bool {
|
|
var target *validationError
|
|
return errors.As(err, &target)
|
|
}
|
|
|
|
// ValidationField 返回表单错误对应的字段名,供页面把焦点放到需要修正的位置。
|
|
func ValidationField(err error) string {
|
|
var target *validationError
|
|
if errors.As(err, &target) {
|
|
return target.field
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func invalidInput(message string, args ...any) error {
|
|
return &validationError{message: fmt.Sprintf(message, args...)}
|
|
}
|
|
|
|
func invalidFieldInput(field, message string, args ...any) error {
|
|
return &validationError{field: field, message: fmt.Sprintf(message, args...)}
|
|
}
|
|
|
|
// 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 err := validateUsername(username); err != nil {
|
|
return err
|
|
}
|
|
hash, err := hashConfirmedPassword(password, confirmation)
|
|
if err != nil {
|
|
return 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,
|
|
})
|
|
}
|
|
|
|
func validateUsername(username string) error {
|
|
if username == "" {
|
|
return invalidInput("用户名不能为空")
|
|
}
|
|
if len([]rune(username)) > maxUsernameLen {
|
|
return invalidInput("用户名不能超过 %d 个字符", maxUsernameLen)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func hashConfirmedPassword(password, confirmation string) ([]byte, error) {
|
|
if len([]rune(password)) < minimumPasswordLen {
|
|
return nil, invalidFieldInput("new_password", "密码至少需要 %d 个字符", minimumPasswordLen)
|
|
}
|
|
// bcrypt 最多接受 72 字节。中文等字符可能占多个字节,因此不能只靠
|
|
// HTML 的 maxlength;服务端需要在哈希前给出可理解的校验错误。
|
|
if len([]byte(password)) > 72 {
|
|
return nil, invalidFieldInput("new_password", "密码不能超过 72 个字节")
|
|
}
|
|
if password != confirmation {
|
|
return nil, invalidFieldInput("password_confirm", "两次输入的密码不一致")
|
|
}
|
|
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("生成密码哈希失败: %w", err)
|
|
}
|
|
return hash, nil
|
|
}
|
|
|
|
// ChangeAdminPassword 校验当前管理员密码,更新新密码并撤销该账号全部 Session。
|
|
func ChangeAdminPassword(db *sql.DB, actor *model.User, currentPassword, newPassword, confirmation string, now time.Time) error {
|
|
if actor == nil || !actor.IsAdmin() {
|
|
return ErrAdminRequired
|
|
}
|
|
if bcrypt.CompareHashAndPassword([]byte(actor.PasswordHash), []byte(currentPassword)) != nil {
|
|
return invalidFieldInput("current_password", "当前密码错误,请重新输入")
|
|
}
|
|
hash, err := hashConfirmedPassword(newPassword, confirmation)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return repository.ChangePasswordAndRevokeSessions(
|
|
db, actor.UserID, actor.PasswordHash, string(hash), now.UTC().Format(model.TimeLayout))
|
|
}
|
|
|
|
// 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
|
|
}
|