Files

231 lines
7.0 KiB
Go

package service
import (
"context"
"errors"
"io"
"log/slog"
"strings"
"testing"
"ws_usernode/internal/auth"
)
// testLogger 丢弃日志输出的测试 logger。
func testLogger() *slog.Logger {
return slog.New(slog.NewTextHandler(io.Discard, nil))
}
// recordingMailer 捕获最近一封邮件,便于从重置邮件中提取令牌。
type recordingMailer struct {
lastTo string
lastSubject string
lastBody string
}
func (m *recordingMailer) Send(_ context.Context, to, subject, body string) error {
m.lastTo = to
m.lastSubject = subject
m.lastBody = body
return nil
}
// newTestAuthService 组装一套完整的认证服务(DB + 内存验证码 + DB OTP/会话/令牌)。
func newTestAuthService(t *testing.T) (*AuthService, *fakeSys, *recordingMailer) {
t.Helper()
db := testDB(t)
sys := newFakeSys()
userSvc := NewUserService(db, sys, testConfig())
adminSvc := NewAdminService(db)
auditSvc := NewAuditService(db, testLogger())
cfg := testConfig()
mailer := &recordingMailer{}
otps := auth.NewDBOTPStore(db, auth.DefaultMaxFailures, auth.DefaultFailureWin)
captchas := auth.NewMemoryCaptchaStore(cfg.Auth.CaptchaTTL)
sessions := auth.NewDBSessionStore(db)
resets := auth.NewDBResetTokenStore(db)
limiter := auth.NewRateLimiter(cfg.Auth.MaxLoginFailures, cfg.Auth.LockDuration)
svc := NewAuthService(db, cfg, otps, captchas, sessions, resets, limiter, mailer, userSvc, adminSvc, auditSvc, testLogger())
return svc, sys, mailer
}
func TestAuthServiceAdminLogin(t *testing.T) {
svc, _, _ := newTestAuthService(t)
ctx := context.Background()
mustAdmin(t, NewAdminService(svc.db))
sid, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test")
if err != nil {
t.Fatalf("admin login: %v", err)
}
if sid == "" {
t.Fatal("session id should not be empty")
}
// 错误密码
if _, err := svc.AdminLogin(ctx, "root", "wrongpass", "127.0.0.1", "test"); !errors.Is(err, ErrBadCredentials) {
t.Fatalf("bad password err = %v, want ErrBadCredentials", err)
}
}
func TestAuthServiceAdminLoginRateLimited(t *testing.T) {
svc, _, _ := newTestAuthService(t)
ctx := context.Background()
mustAdmin(t, NewAdminService(svc.db))
for i := 0; i < 5; i++ {
_, _ = svc.AdminLogin(ctx, "root", "wrongpass", "127.0.0.1", "test")
}
_, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test")
if !errors.Is(err, ErrRateLimited) {
t.Fatalf("login after lock err = %v, want ErrRateLimited", err)
}
}
func TestAuthServiceUserOTPFlow(t *testing.T) {
svc, sys, mailer := newTestAuthService(t)
ctx := context.Background()
// 准备一个外部用户
userSvc := NewUserService(svc.db, sys, testConfig())
u, err := userSvc.Create(ctx, "zhangsan", "zs@example.com", "", "", 0, 0)
if err != nil {
t.Fatalf("create user: %v", err)
}
// 获取图形验证码
cap, err := svc.captchas.New()
if err != nil {
t.Fatalf("captcha: %v", err)
}
// 发送 OTP
if err := svc.UserOTPSend(ctx, u.Username, cap.ID, cap.Text, "127.0.0.1"); err != nil {
t.Fatalf("otp send: %v", err)
}
if !strings.Contains(mailer.lastBody, "验证码") {
t.Fatalf("otp mail body unexpected: %s", mailer.lastBody)
}
// CLI 通道复用同一验证码
code, err := svc.otps.Current(ctx, u.Username)
if err != nil {
t.Fatalf("otp current: %v", err)
}
if len(code) != 6 {
t.Fatalf("otp len = %d, want 6", len(code))
}
// OTP 登录
sid, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test")
if err != nil {
t.Fatalf("otp login: %v", err)
}
if sid == "" {
t.Fatal("session id should not be empty")
}
// 一次性:再次使用失败
if _, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test"); err != auth.ErrInvalidCode {
t.Fatalf("reuse otp err = %v, want ErrInvalidCode", err)
}
// 验证码错误时 send 被拒绝
cap2, _ := svc.captchas.New()
if err := svc.UserOTPSend(ctx, u.Username, cap2.ID, "0000", "127.0.0.1"); !errors.Is(err, ErrCaptchaFailed) {
t.Fatalf("bad captcha err = %v, want ErrCaptchaFailed", err)
}
// 冷却期内再次 send 被拒
cap3, _ := svc.captchas.New()
if err := svc.UserOTPSend(ctx, u.Username, cap3.ID, cap3.Text, "127.0.0.1"); !errors.Is(err, auth.ErrCooldown) {
t.Fatalf("cooldown err = %v, want ErrCooldown", err)
}
}
func TestAuthServiceAdminForgotReset(t *testing.T) {
svc, _, mailer := newTestAuthService(t)
ctx := context.Background()
adm := mustAdmin(t, NewAdminService(svc.db))
// forgot → 邮件应包含重置链接
if err := svc.AdminForgot(ctx, adm.Username, "127.0.0.1"); err != nil {
t.Fatalf("forgot: %v", err)
}
if mailer.lastTo != adm.Email {
t.Fatalf("mail to = %q, want %q", mailer.lastTo, adm.Email)
}
// 从邮件 body 提取 token
idx := strings.Index(mailer.lastBody, "token=")
if idx < 0 {
t.Fatalf("reset link missing token: %s", mailer.lastBody)
}
token := mailer.lastBody[idx+len("token="):]
token = strings.TrimSpace(strings.SplitN(token, "\n", 2)[0])
// reset
if err := svc.AdminReset(ctx, token, "NewPassw0rd", "127.0.0.1"); err != nil {
t.Fatalf("reset: %v", err)
}
// 新密码可登录
adminSvc := NewAdminService(svc.db)
if _, err := adminSvc.Login(ctx, adm.Username, "NewPassw0rd"); err != nil {
t.Fatalf("login with new password: %v", err)
}
// 旧密码失效
if _, err := adminSvc.Login(ctx, adm.Username, "Passw0rd"); !errors.Is(err, ErrBadCredentials) {
t.Fatalf("login with old password err = %v, want ErrBadCredentials", err)
}
// 令牌一次性
if err := svc.AdminReset(ctx, token, "AgainPassw0rd", "127.0.0.1"); !errors.Is(err, auth.ErrResetTokenInvalid) {
t.Fatalf("reuse token err = %v, want ErrResetTokenInvalid", err)
}
// 不存在的用户 forgot 也应成功(防枚举)
if err := svc.AdminForgot(ctx, "ghost", "127.0.0.1"); err != nil {
t.Fatalf("forgot ghost: %v", err)
}
}
func TestAuthServiceMeLogout(t *testing.T) {
svc, sys, _ := newTestAuthService(t)
ctx := context.Background()
mustAdmin(t, NewAdminService(svc.db))
sid, err := svc.AdminLogin(ctx, "root", "Passw0rd", "127.0.0.1", "test")
if err != nil {
t.Fatalf("login: %v", err)
}
info, err := svc.Me(ctx, sid)
if err != nil {
t.Fatalf("me: %v", err)
}
if info.UserType != "admin" || info.Username != "root" {
t.Fatalf("me info = %+v", info)
}
if err := svc.Logout(ctx, sid); err != nil {
t.Fatalf("logout: %v", err)
}
if _, err := svc.Me(ctx, sid); !errors.Is(err, auth.ErrSessionNotFound) {
t.Fatalf("me after logout err = %v, want ErrSessionNotFound", err)
}
// 外部用户 me
userSvc := NewUserService(svc.db, sys, testConfig())
u, _ := userSvc.Create(ctx, "lisi", "ls@example.com", "", "", 0, 0)
cap, _ := svc.captchas.New()
_ = svc.UserOTPSend(ctx, u.Username, cap.ID, cap.Text, "127.0.0.1")
code, _ := svc.otps.Current(ctx, u.Username)
usid, err := svc.UserOTPLogin(ctx, u.Username, code, "127.0.0.1", "test")
if err != nil {
t.Fatalf("user login: %v", err)
}
info, err = svc.Me(ctx, usid)
if err != nil {
t.Fatalf("user me: %v", err)
}
if info.UserType != "user" || info.Username != "ext_lisi" {
t.Fatalf("user me info = %+v", info)
}
}