feat(M1): 认证与用户管理 — 双通道登录、cookie 会话、用户 CRUD 与真实系统账号对接

认证:
- 图形验证码 GET /auth/captcha(内置 PNG 渲染,零第三方依赖)
- 外部用户 OTP 双通道:DB 存储(otp_codes)使邮件与 CLI 共用同一验证码/冷却/失败限速
- 管理员 bcrypt 登录 + 连续失败限速锁定;admin/forgot + admin/reset 邮件重置(SMTP 或日志)
- cookie 会话(HttpOnly/SameSite)、me/logout、admin/user 鉴权中间件

用户管理(admin):
- CRUD + disable/enable/extend/delete,对接 system 层真实 useradd/usermod/userdel/passwd
- system 层三执行模式:dry-run(默认,安全)/ direct(容器/测试用户)/ sudo(生产 sudoers 白名单)
- Exists 系统账号一致性检查;deploy/sudoers.example 白名单模板
- 关键操作接入 append-only 审计

其他:
- CLI user otp 改 DB store,与邮件通道真正对齐
- 容器镜像补 shadow(alpine 无 useradd);Makefile VERSION 0.2.0-m1
- 测试:auth/service 单测 + api httptest 集成 + 容器内真实系统账号端到端验证
This commit is contained in:
2026-08-29 23:40:20 +08:00
parent ae45aba607
commit 630d240dc0
32 changed files with 2923 additions and 188 deletions
+7 -5
View File
@@ -14,10 +14,11 @@ var ErrCaptchaInvalid = errors.New("auth: 图形验证码错误")
// Captcha 图形验证码(防机器人,登录前置)。
type Captcha struct {
ID string
Text string // M1 生成图像渲染,此处仅存文本
Text string
}
// CaptchaStore 为图形验证码存储M1 实现图像渲染)。
// CaptchaStore 为图形验证码存储。单实例内存实现为默认;
// 多实例部署需改 DB/RedisPLAN §6 注明)。
type CaptchaStore interface {
// New 生成一个验证码并返回其 ID。
New() (*Captcha, error)
@@ -27,6 +28,7 @@ type CaptchaStore interface {
// MemoryCaptchaStore 单实例内存实现。
type MemoryCaptchaStore struct {
ttl time.Duration
mu sync.Mutex
entries map[string]*captchaEntry
}
@@ -37,8 +39,8 @@ type captchaEntry struct {
}
// NewMemoryCaptchaStore 创建内存图形验证码存储。
func NewMemoryCaptchaStore() *MemoryCaptchaStore {
return &MemoryCaptchaStore{entries: make(map[string]*captchaEntry)}
func NewMemoryCaptchaStore(ttl time.Duration) *MemoryCaptchaStore {
return &MemoryCaptchaStore{ttl: ttl, entries: make(map[string]*captchaEntry)}
}
func (s *MemoryCaptchaStore) New() (*Captcha, error) {
@@ -52,7 +54,7 @@ func (s *MemoryCaptchaStore) New() (*Captcha, error) {
}
s.mu.Lock()
defer s.mu.Unlock()
s.entries[id] = &captchaEntry{text: text, expiresAt: time.Now().Add(5 * time.Minute)}
s.entries[id] = &captchaEntry{text: text, expiresAt: time.Now().Add(s.ttl)}
return &Captcha{ID: id, Text: text}, nil
}
+131
View File
@@ -0,0 +1,131 @@
package auth
import (
"bytes"
"fmt"
"image"
"image/color"
"image/draw"
"image/png"
"math/rand"
"time"
)
// 验证码图像渲染:内置 5x7 点阵数字 + 噪点 + 干扰线,不依赖第三方字体库
// (PLAN §4:内置生成,不依赖第三方服务)。
// digitGlyphs 为 0-9 的 5x7 点阵,每行低 5 位表示该行像素。
var digitGlyphs = [10][7]byte{
{0b01110, 0b10001, 0b10011, 0b10101, 0b11001, 0b10001, 0b01110}, // 0
{0b00100, 0b01100, 0b00100, 0b00100, 0b00100, 0b00100, 0b01110}, // 1
{0b01110, 0b10001, 0b00001, 0b00010, 0b00100, 0b01000, 0b11111}, // 2
{0b11111, 0b00010, 0b00100, 0b00010, 0b00001, 0b10001, 0b01110}, // 3
{0b00010, 0b00110, 0b01010, 0b10010, 0b11111, 0b00010, 0b00010}, // 4
{0b11111, 0b10000, 0b11110, 0b00001, 0b00001, 0b10001, 0b01110}, // 5
{0b00110, 0b01000, 0b10000, 0b11110, 0b10001, 0b10001, 0b01110}, // 6
{0b11111, 0b00001, 0b00010, 0b00100, 0b01000, 0b01000, 0b01000}, // 7
{0b01110, 0b10001, 0b10001, 0b01110, 0b10001, 0b10001, 0b01110}, // 8
{0b01110, 0b10001, 0b10001, 0b01111, 0b00001, 0b00010, 0b01100}, // 9
}
const (
glyphW = 5
glyphH = 7
scale = 3 // 点阵放大倍数
charGap = 4 // 字符间距(像素)
edgePad = 6 // 画布边距
)
// RenderCaptchaPNG 将 4 位数字验证码渲染为 PNG 字节流。
// 仅接受数字字符,其余返回错误。
func RenderCaptchaPNG(text string) ([]byte, error) {
if len(text) == 0 || len(text) > 8 {
return nil, fmt.Errorf("auth: captcha text length must be 1~8")
}
for _, r := range text {
if r < '0' || r > '9' {
return nil, fmt.Errorf("auth: captcha text must be digits")
}
}
canvasW := edgePad*2 + len(text)*glyphW*scale + (len(text)-1)*charGap
canvasH := edgePad*2 + glyphH*scale
img := image.NewRGBA(image.Rect(0, 0, canvasW, canvasH))
draw.Draw(img, img.Bounds(), &image.Uniform{C: color.RGBA{248, 250, 252, 255}}, image.Point{}, draw.Src)
rng := rand.New(rand.NewSource(time.Now().UnixNano()))
// 噪点(浅灰,稀疏)
for i := 0; i < canvasW*canvasH/7; i++ {
x, y := rng.Intn(canvasW), rng.Intn(canvasH)
g := uint8(150 + rng.Intn(90))
img.Set(x, y, color.RGBA{g, g, g, 255})
}
// 干扰线(穿过字符区域的浅色斜线)
for i := 0; i < 3; i++ {
g := uint8(160 + rng.Intn(80))
c := color.RGBA{g, g, g, 255}
x1, y1 := rng.Intn(canvasW/2), rng.Intn(canvasH)
x2, y2 := canvasW/2+rng.Intn(canvasW/2), rng.Intn(canvasH)
drawLine(img, x1, y1, x2, y2, c)
}
// 逐字符绘制(颜色随机取深色系)
inkPalette := []color.RGBA{
{40, 60, 110, 255}, {120, 45, 45, 255}, {30, 90, 60, 255}, {80, 60, 110, 255},
}
for i, r := range text {
glyph := digitGlyphs[r-'0']
ink := inkPalette[rng.Intn(len(inkPalette))]
x0 := edgePad + i*(glyphW*scale+charGap)
y0 := edgePad
for row := 0; row < glyphH; row++ {
for col := 0; col < glyphW; col++ {
if glyph[row]&(1<<(glyphW-1-col)) != 0 {
fillRect(img, x0+col*scale, y0+row*scale, scale, scale, ink)
}
}
}
}
var buf bytes.Buffer
if err := png.Encode(&buf, img); err != nil {
return nil, fmt.Errorf("auth: encode captcha png: %w", err)
}
return buf.Bytes(), nil
}
// fillRect 填充实心矩形。
func fillRect(img *image.RGBA, x, y, w, h int, c color.RGBA) {
for dy := 0; dy < h; dy++ {
for dx := 0; dx < w; dx++ {
img.Set(x+dx, y+dy, c)
}
}
}
// drawLine 使用 DDA 算法画线。
func drawLine(img *image.RGBA, x1, y1, x2, y2 int, c color.RGBA) {
steps := abs(x2-x1)
if d := abs(y2 - y1); d > steps {
steps = d
}
if steps == 0 {
img.Set(x1, y1, c)
return
}
for i := 0; i <= steps; i++ {
x := x1 + (x2-x1)*i/steps
y := y1 + (y2-y1)*i/steps
if x >= 0 && x < img.Bounds().Dx() && y >= 0 && y < img.Bounds().Dy() {
img.Set(x, y, c)
}
}
}
func abs(n int) int {
if n < 0 {
return -n
}
return n
}
+115
View File
@@ -0,0 +1,115 @@
package auth
import (
"context"
"image/png"
"strings"
"testing"
"time"
)
func TestRenderCaptchaPNG(t *testing.T) {
pngBytes, err := RenderCaptchaPNG("4837")
if err != nil {
t.Fatalf("render: %v", err)
}
img, err := png.Decode(strings.NewReader(string(pngBytes)))
if err != nil {
t.Fatalf("decode png: %v", err)
}
if img.Bounds().Dx() < 50 || img.Bounds().Dy() < 20 {
t.Fatalf("canvas too small: %dx%d", img.Bounds().Dx(), img.Bounds().Dy())
}
}
func TestRenderCaptchaPNGInvalid(t *testing.T) {
if _, err := RenderCaptchaPNG("12a4"); err == nil {
t.Fatal("expected error for non-digit text")
}
if _, err := RenderCaptchaPNG(""); err == nil {
t.Fatal("expected error for empty text")
}
}
func TestMemoryCaptchaStore(t *testing.T) {
s := NewMemoryCaptchaStore(5 * time.Minute)
cap, err := s.New()
if err != nil {
t.Fatalf("new: %v", err)
}
if len(cap.Text) != 4 {
t.Fatalf("captcha text len = %d, want 4", len(cap.Text))
}
if !s.Verify(cap.ID, cap.Text) {
t.Fatal("verify should succeed")
}
if s.Verify(cap.ID, cap.Text) {
t.Fatal("captcha must be one-time")
}
}
func TestMemoryCaptchaStoreExpired(t *testing.T) {
s := NewMemoryCaptchaStore(-time.Second) // 立即过期
cap, err := s.New()
if err != nil {
t.Fatalf("new: %v", err)
}
if s.Verify(cap.ID, cap.Text) {
t.Fatal("expired captcha should fail")
}
}
func TestRateLimiter(t *testing.T) {
l := NewRateLimiter(3, time.Minute)
key := "admin-login:admin"
if !l.Allow(key) {
t.Fatal("first attempt should be allowed")
}
for i := 0; i < 3; i++ {
l.RecordFailure(key)
}
if l.Allow(key) {
t.Fatal("should be locked after max failures")
}
l.Reset(key)
if !l.Allow(key) {
t.Fatal("should be allowed after reset")
}
}
func TestDBResetTokenStore(t *testing.T) {
db := testDB(t)
s := NewDBResetTokenStore(db)
ctx := context.Background()
token, err := s.Create(ctx, 42, time.Minute, "127.0.0.1")
if err != nil {
t.Fatalf("create: %v", err)
}
if token == "" {
t.Fatal("token should not be empty")
}
adminID, err := s.Consume(ctx, token)
if err != nil {
t.Fatalf("consume: %v", err)
}
if adminID != 42 {
t.Fatalf("adminID = %d, want 42", adminID)
}
// 一次性
if _, err := s.Consume(ctx, token); err != ErrResetTokenInvalid {
t.Fatalf("second consume err = %v, want ErrResetTokenInvalid", err)
}
// 无效令牌
if _, err := s.Consume(ctx, "not-a-token"); err != ErrResetTokenInvalid {
t.Fatalf("bad token err = %v, want ErrResetTokenInvalid", err)
}
// 过期令牌
dbToken, err := s.Create(ctx, 7, -time.Minute, "")
if err != nil {
t.Fatalf("create expired: %v", err)
}
if _, err := s.Consume(ctx, dbToken); err != ErrResetTokenInvalid {
t.Fatalf("expired token err = %v, want ErrResetTokenInvalid", err)
}
}
+159 -17
View File
@@ -1,15 +1,22 @@
// Package auth 提供认证相关能力:OTP(双通道)、会话、bcrypt、图形验证码。
//
// OTP 双通道对齐:邮件发送与 CLI 获取共用同一 OTPStore(同一验证码、同一
// 10 分钟有效期、同一 60s 冷却与失败限速),邮件失败不阻断 CLI 通道。
// 单实例用内存存储;多实例需改为 DB/Redis(PLAN §6 注明)。
// 有效期、同一冷却与失败限速),邮件失败不阻断 CLI 通道。邮件通道经
// Send 生成验证码,CLI 通道优先经 Current 复用同一验证码,无有效码时才
// 触发 Send(仍受同一冷却约束)。验证码落 DB(otp_codes 表),单实例部署
// 即可保证跨进程(HTTP 服务与 CLI 子命令)共享;多实例需改 DB 行锁/Redis
// PLAN §6)。
package auth
import (
"context"
"errors"
"sync"
"time"
"gorm.io/gorm"
"ws_usernode/internal/model"
"ws_usernode/internal/pkg"
)
@@ -21,20 +28,141 @@ var (
)
const (
otpCodeLen = 6
maxFailures = 5 // 单账号连续失败限速阈值
failureWin = 10 * time.Minute // 失败计数窗口
otpCodeLen = 6
)
// OTPStore 为 OTP 验证码存储。内存实现为单实例默认实现
// OTP 失败限速默认参数(单账号连续失败阈值与计数窗口)
const (
DefaultMaxFailures = 5
DefaultFailureWin = 10 * time.Minute
)
// OTPStore 为 OTP 验证码存储。DB 实现(DBOTPStore)为生产默认,
// MemoryOTPStore 供测试与单进程内嵌场景使用。
type OTPStore interface {
// Send 为 username 生成新验证码覆盖旧码)。冷却期内调用返回 ErrCooldown。
// 邮件与 CLI 双通道都走该方法,保证对齐。
Send(username string, ttl, cooldown time.Duration) (string, error)
// Send 为 username 生成新验证码覆盖旧码(邮件通道)。冷却期内返回 ErrCooldown。
Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error)
// Current 返回当前有效(未过期、未消费)验证码,供 CLI 通道复用同一验证码。
// 无有效验证码返回 ErrInvalidCode。
Current(ctx context.Context, username string) (string, error)
// Verify 校验验证码并一次性消费。失败累计计数(达到阈值返回 ErrTooManyFails)。
Verify(username, code string) (bool, error)
Verify(ctx context.Context, username, code string) (bool, error)
// Failures 返回 username 当前失败计数。
Failures(username string) (int, error)
Failures(ctx context.Context, username string) (int, error)
}
// DBOTPStore 基于 model.OTPCode 的存储实现,每用户一行(username 唯一)。
type DBOTPStore struct {
db *gorm.DB
maxFailures int
failureWin time.Duration
}
// NewDBOTPStore 创建 DB OTP 存储。
func NewDBOTPStore(db *gorm.DB, maxFailures int, failureWin time.Duration) *DBOTPStore {
return &DBOTPStore{db: db, maxFailures: maxFailures, failureWin: failureWin}
}
func (s *DBOTPStore) Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error) {
now := time.Now()
var row model.OTPCode
err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error
// 注意:First 的 err 不能复用给后续语句,避免被覆盖导致走错分支
exists := err == nil
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return "", err
}
if exists && now.Before(row.CooldownUntil) {
return "", ErrCooldown
}
code, randErr := pkg.RandomDigits(otpCodeLen)
if randErr != nil {
return "", randErr
}
updates := map[string]any{
"code": code,
"expires_at": now.Add(ttl),
"cooldown_until": now.Add(cooldown),
"failures": 0,
"failed_at": nil,
"consumed_at": nil,
}
if exists {
err = s.db.WithContext(ctx).Model(&row).Updates(updates).Error
} else {
err = s.db.WithContext(ctx).Create(&model.OTPCode{
Username: username,
Code: code,
ExpiresAt: now.Add(ttl),
CooldownUntil: now.Add(cooldown),
}).Error
}
if err != nil {
return "", err
}
return code, nil
}
func (s *DBOTPStore) Current(ctx context.Context, username string) (string, error) {
var row model.OTPCode
if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", ErrInvalidCode
}
return "", err
}
if row.ConsumedAt != nil || time.Now().After(row.ExpiresAt) {
return "", ErrInvalidCode
}
return row.Code, nil
}
func (s *DBOTPStore) Verify(ctx context.Context, username, code string) (bool, error) {
now := time.Now()
var row model.OTPCode
if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, ErrInvalidCode
}
return false, err
}
if row.ConsumedAt != nil || now.After(row.ExpiresAt) {
return false, ErrInvalidCode
}
// 失败计数窗口:超出窗口则重置计数
if row.FailedAt != nil && now.Sub(*row.FailedAt) > s.failureWin {
row.Failures = 0
row.FailedAt = nil
}
if row.Failures >= s.maxFailures {
return false, ErrTooManyFails
}
if row.Code != code {
row.Failures++
f := now
row.FailedAt = &f
_ = s.db.WithContext(ctx).Model(&row).Updates(map[string]any{"failures": row.Failures, "failed_at": row.FailedAt}).Error
if row.Failures >= s.maxFailures {
return false, ErrTooManyFails
}
return false, ErrInvalidCode
}
consumed := now
if err := s.db.WithContext(ctx).Model(&row).Update("consumed_at", &consumed).Error; err != nil {
return false, err
}
return true, nil
}
func (s *DBOTPStore) Failures(ctx context.Context, username string) (int, error) {
var row model.OTPCode
if err := s.db.WithContext(ctx).Where("username = ?", username).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return 0, nil
}
return 0, err
}
return row.Failures, nil
}
type otpEntry struct {
@@ -44,7 +172,7 @@ type otpEntry struct {
failures int
}
// MemoryOTPStore 为单实例内存实现。
// MemoryOTPStore 为单进程内存实现(测试/内嵌场景)
type MemoryOTPStore struct {
mu sync.Mutex
entries map[string]*otpEntry
@@ -55,9 +183,10 @@ func NewMemoryOTPStore() *MemoryOTPStore {
return &MemoryOTPStore{entries: make(map[string]*otpEntry)}
}
func (s *MemoryOTPStore) Send(username string, ttl, cooldown time.Duration) (string, error) {
func (s *MemoryOTPStore) Send(ctx context.Context, username string, ttl, cooldown time.Duration) (string, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
now := time.Now()
if e, ok := s.entries[username]; ok && now.Before(e.cooldownAt) {
@@ -76,9 +205,21 @@ func (s *MemoryOTPStore) Send(username string, ttl, cooldown time.Duration) (str
return code, nil
}
func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
func (s *MemoryOTPStore) Current(ctx context.Context, username string) (string, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
e, ok := s.entries[username]
if !ok || time.Now().After(e.expiresAt) {
return "", ErrInvalidCode
}
return e.code, nil
}
func (s *MemoryOTPStore) Verify(ctx context.Context, username, code string) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
e, ok := s.entries[username]
if !ok {
@@ -89,12 +230,12 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
delete(s.entries, username)
return false, ErrInvalidCode
}
if e.failures >= maxFailures {
if e.failures >= DefaultMaxFailures {
return false, ErrTooManyFails
}
if e.code != code {
e.failures++
if e.failures >= maxFailures {
if e.failures >= DefaultMaxFailures {
return false, ErrTooManyFails
}
return false, ErrInvalidCode
@@ -103,9 +244,10 @@ func (s *MemoryOTPStore) Verify(username, code string) (bool, error) {
return true, nil
}
func (s *MemoryOTPStore) Failures(username string) (int, error) {
func (s *MemoryOTPStore) Failures(ctx context.Context, username string) (int, error) {
s.mu.Lock()
defer s.mu.Unlock()
_ = ctx
if e, ok := s.entries[username]; ok {
return e.failures, nil
}
+119
View File
@@ -0,0 +1,119 @@
package auth
import (
"context"
"testing"
"time"
"gorm.io/gorm"
"ws_usernode/internal/model"
)
func testDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := model.Open("sqlite", ":memory:", false)
if err != nil {
t.Fatalf("open test db: %v", err)
}
if err := model.Migrate(db); err != nil {
t.Fatalf("migrate: %v", err)
}
return db
}
func TestDBOTPStoreSendCurrentVerify(t *testing.T) {
db := testDB(t)
s := NewDBOTPStore(db, DefaultMaxFailures, DefaultFailureWin)
ctx := context.Background()
const user = "ext_zhangsan"
code, err := s.Send(ctx, user, 10*time.Minute, time.Minute)
if err != nil {
t.Fatalf("send: %v", err)
}
if len(code) != 6 {
t.Fatalf("code length = %d, want 6", len(code))
}
// 双通道对齐:Current 复用同一验证码
cur, err := s.Current(ctx, user)
if err != nil {
t.Fatalf("current: %v", err)
}
if cur != code {
t.Fatalf("current = %q, want %q (双通道必须同一验证码)", cur, code)
}
// 校验成功
ok, err := s.Verify(ctx, user, code)
if err != nil || !ok {
t.Fatalf("verify = %v/%v, want true/nil", ok, err)
}
// 一次性:再次校验失败
ok, err = s.Verify(ctx, user, code)
if err != ErrInvalidCode {
t.Fatalf("second verify err = %v, want ErrInvalidCode", err)
}
if ok {
t.Fatal("second verify should fail")
}
// Current 也应失败(已消费)
if _, err := s.Current(ctx, user); err != ErrInvalidCode {
t.Fatalf("current after consume err = %v, want ErrInvalidCode", err)
}
}
func TestDBOTPStoreCooldown(t *testing.T) {
db := testDB(t)
s := NewDBOTPStore(db, DefaultMaxFailures, DefaultFailureWin)
ctx := context.Background()
if _, err := s.Send(ctx, "ext_lisi", time.Minute, time.Minute); err != nil {
t.Fatalf("send: %v", err)
}
if _, err := s.Send(ctx, "ext_lisi", time.Minute, time.Minute); err != ErrCooldown {
t.Fatalf("second send err = %v, want ErrCooldown", err)
}
}
func TestDBOTPStoreFailuresAndWindow(t *testing.T) {
db := testDB(t)
s := NewDBOTPStore(db, 3, 10*time.Minute)
ctx := context.Background()
code, err := s.Send(ctx, "ext_wangwu", time.Minute, 0)
if err != nil {
t.Fatalf("send: %v", err)
}
// 3 次错误后进入限速
for i := 0; i < 3; i++ {
if _, err := s.Verify(ctx, "ext_wangwu", "000000"); err != ErrInvalidCode && err != ErrTooManyFails {
t.Fatalf("verify(%d) err = %v", i, err)
}
}
if _, err := s.Verify(ctx, "ext_wangwu", code); err != ErrTooManyFails {
t.Fatalf("verify after max failures err = %v, want ErrTooManyFails", err)
}
f, err := s.Failures(ctx, "ext_wangwu")
if err != nil {
t.Fatalf("failures: %v", err)
}
if f != 3 {
t.Fatalf("failures = %d, want 3", f)
}
}
func TestMemoryOTPStore(t *testing.T) {
s := NewMemoryOTPStore()
ctx := context.Background()
code, err := s.Send(ctx, "ext_test", time.Minute, time.Second)
if err != nil {
t.Fatalf("send: %v", err)
}
if cur, err := s.Current(ctx, "ext_test"); err != nil || cur != code {
t.Fatalf("current = %q/%v, want %q/nil", cur, err, code)
}
if _, err := s.Send(ctx, "ext_test", time.Minute, time.Second); err != ErrCooldown {
t.Fatalf("send during cooldown err = %v, want ErrCooldown", err)
}
if ok, err := s.Verify(ctx, "ext_test", code); err != nil || !ok {
t.Fatalf("verify = %v/%v, want true/nil", ok, err)
}
}
+68
View File
@@ -0,0 +1,68 @@
package auth
import (
"sync"
"time"
)
// RateLimiter 内存登录限速器:按 key(如用户名)统计连续失败次数,
// 达到阈值后锁定 lockFor 时长,锁定期内 Allow 返回 false。
// 单实例内存实现即可满足(登录失败限速无跨进程一致性要求)。
type RateLimiter struct {
mu sync.Mutex
max int
lockFor time.Duration
entries map[string]*rlEntry
now func() time.Time // 可注入时钟(测试)
}
type rlEntry struct {
failures int
lockedUntil time.Time
}
// NewRateLimiter 创建限速器。
func NewRateLimiter(max int, lockFor time.Duration) *RateLimiter {
return &RateLimiter{max: max, lockFor: lockFor, entries: make(map[string]*rlEntry), now: time.Now}
}
// Allow 返回 key 当前是否允许继续尝试。
func (l *RateLimiter) Allow(key string) bool {
l.mu.Lock()
defer l.mu.Unlock()
e, ok := l.entries[key]
if !ok {
return true
}
if now := l.now(); now.Before(e.lockedUntil) {
return false
} else if e.failures >= l.max {
// 锁定已过期:重置计数,允许重试
delete(l.entries, key)
return true
}
return true
}
// RecordFailure 记录一次失败;达到阈值后进入锁定。
func (l *RateLimiter) RecordFailure(key string) {
l.mu.Lock()
defer l.mu.Unlock()
now := l.now()
e, ok := l.entries[key]
if !ok || now.After(e.lockedUntil) && e.failures >= l.max {
e = &rlEntry{}
l.entries[key] = e
}
e.failures++
if e.failures >= l.max {
e.lockedUntil = now.Add(l.lockFor)
}
}
// Reset 清除 key 的失败记录(登录成功后调用)。
func (l *RateLimiter) Reset(key string) {
l.mu.Lock()
defer l.mu.Unlock()
delete(l.entries, key)
}
+76
View File
@@ -0,0 +1,76 @@
package auth
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"time"
"gorm.io/gorm"
"ws_usernode/internal/model"
"ws_usernode/internal/pkg"
)
// ErrResetTokenInvalid 表示重置令牌无效、已使用或已过期。
var ErrResetTokenInvalid = errors.New("auth: 重置令牌无效或已过期")
// ResetTokenStore 为管理员密码重置令牌存储(邮件重置)。
type ResetTokenStore interface {
// Create 生成令牌并存储其哈希,返回令牌明文(仅经邮件/日志发出)。
Create(ctx context.Context, adminID uint, ttl time.Duration, ip string) (string, error)
// Consume 校验令牌并标记已使用,返回对应的管理员 ID。
Consume(ctx context.Context, token string) (uint, error)
}
// DBResetTokenStore 基于 model.PasswordResetToken 的存储实现。
type DBResetTokenStore struct {
db *gorm.DB
}
// NewDBResetTokenStore 创建重置令牌存储。
func NewDBResetTokenStore(db *gorm.DB) *DBResetTokenStore {
return &DBResetTokenStore{db: db}
}
func (s *DBResetTokenStore) Create(ctx context.Context, adminID uint, ttl time.Duration, ip string) (string, error) {
token, err := pkg.RandomHex(24)
if err != nil {
return "", err
}
row := model.PasswordResetToken{
AdminID: adminID,
TokenHash: hashToken(token),
ExpiresAt: time.Now().Add(ttl),
IP: ip,
}
if err := s.db.WithContext(ctx).Create(&row).Error; err != nil {
return "", err
}
return token, nil
}
func (s *DBResetTokenStore) Consume(ctx context.Context, token string) (uint, error) {
var row model.PasswordResetToken
if err := s.db.WithContext(ctx).Where("token_hash = ?", hashToken(token)).First(&row).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return 0, ErrResetTokenInvalid
}
return 0, err
}
if row.UsedAt != nil || time.Now().After(row.ExpiresAt) {
return 0, ErrResetTokenInvalid
}
now := time.Now()
if err := s.db.WithContext(ctx).Model(&row).Update("used_at", &now).Error; err != nil {
return 0, err
}
return row.AdminID, nil
}
// hashToken 计算令牌的 SHA-256 摘要(令牌本身为高熵随机串,无需加盐)。
func hashToken(token string) string {
sum := sha256.Sum256([]byte(token))
return hex.EncodeToString(sum[:])
}