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) 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) } }