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:
@@ -14,10 +14,11 @@ var ErrCaptchaInvalid = errors.New("auth: 图形验证码错误")
|
||||
// Captcha 图形验证码(防机器人,登录前置)。
|
||||
type Captcha struct {
|
||||
ID string
|
||||
Text string // M1 生成图像渲染,此处仅存文本
|
||||
Text string
|
||||
}
|
||||
|
||||
// CaptchaStore 为图形验证码存储(M1 实现图像渲染)。
|
||||
// CaptchaStore 为图形验证码存储。单实例内存实现为默认;
|
||||
// 多实例部署需改 DB/Redis(PLAN §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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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[:])
|
||||
}
|
||||
Reference in New Issue
Block a user