package twofa import ( "context" "errors" "testing" "github.com/enterprise-ai-platform/server/pkg/auth" ) type fakeStore struct { enabled bool remaining int secret string saved bool enabled2 bool // Enable 被调用 disabled bool // Disable 被调用 backupOK bool getErr error } func (f *fakeStore) GetStatus(ctx context.Context, userID string) (bool, int, error) { return f.enabled, f.remaining, f.getErr } func (f *fakeStore) GetSecret(ctx context.Context, userID string) (string, error) { return f.secret, f.getErr } func (f *fakeStore) GetSecretAndEnabled(ctx context.Context, userID string) (string, bool, error) { return f.secret, f.enabled, f.getErr } func (f *fakeStore) SaveEnrollment(ctx context.Context, userID, secret string, codeHashes []string) error { f.saved = true f.secret = secret return nil } func (f *fakeStore) Enable(ctx context.Context, userID string) error { f.enabled2 = true; return nil } func (f *fakeStore) Disable(ctx context.Context, userID string) error { f.disabled = true; return nil } func (f *fakeStore) ConsumeBackupCode(ctx context.Context, userID, code string) (bool, error) { return f.backupOK, nil } const fixedNow int64 = 1_700_000_000 func newSvc(store Store) *Service { s := NewService(store) s.nowUnix = func() int64 { return fixedNow } return s } func TestEnroll_RejectsWhenAlreadyEnabled(t *testing.T) { svc := newSvc(&fakeStore{enabled: true}) if _, err := svc.Enroll(context.Background(), "u1", "a@b.c"); !errors.Is(err, ErrAlreadyEnabled) { t.Fatalf("已启用应返回 ErrAlreadyEnabled,实际: %v", err) } } func TestEnroll_GeneratesAndSaves(t *testing.T) { store := &fakeStore{} svc := newSvc(store) res, err := svc.Enroll(context.Background(), "u1", "admin@govai.gov.cn") if err != nil { t.Fatalf("Enroll 出错: %v", err) } if res.Secret == "" || len(res.BackupCodes) != 8 { t.Fatalf("应返回密钥与 8 个备份码,实际 codes=%d", len(res.BackupCodes)) } if !store.saved { t.Fatal("应调用 SaveEnrollment") } if res.OtpauthURI == "" { t.Fatal("应返回 otpauth URI") } } func TestEnableAfterVerify(t *testing.T) { secret, _ := auth.GenerateTOTPSecret() code, _ := auth.TOTPCodeAt(secret, fixedNow) // 未设置密钥 if err := newSvc(&fakeStore{secret: ""}).EnableAfterVerify(context.Background(), "u1", code); !errors.Is(err, ErrNotSetup) { t.Fatalf("无密钥应返回 ErrNotSetup,实际: %v", err) } // 错误验证码 if err := newSvc(&fakeStore{secret: secret}).EnableAfterVerify(context.Background(), "u1", "000000"); !errors.Is(err, ErrBadCode) { t.Fatalf("错误码应返回 ErrBadCode,实际: %v", err) } // 正确验证码 store := &fakeStore{secret: secret} if err := newSvc(store).EnableAfterVerify(context.Background(), "u1", code); err != nil { t.Fatalf("正确码应成功,实际: %v", err) } if !store.enabled2 { t.Fatal("应调用 Enable") } } func TestDisable(t *testing.T) { secret, _ := auth.GenerateTOTPSecret() code, _ := auth.TOTPCodeAt(secret, fixedNow) // 未启用 → 幂等成功,不调用 Disable store0 := &fakeStore{enabled: false} if err := newSvc(store0).Disable(context.Background(), "u1", "", ""); err != nil || store0.disabled { t.Fatalf("未启用应幂等返回 nil 且不调用 Disable,err=%v disabled=%v", err, store0.disabled) } // 启用 + 正确 TOTP store1 := &fakeStore{enabled: true, secret: secret} if err := newSvc(store1).Disable(context.Background(), "u1", code, ""); err != nil { t.Fatalf("正确 TOTP 应成功: %v", err) } if !store1.disabled { t.Fatal("应调用 Disable") } // 启用 + 备份码 store2 := &fakeStore{enabled: true, secret: secret, backupOK: true} if err := newSvc(store2).Disable(context.Background(), "u1", "", "backup-xxxx"); err != nil || !store2.disabled { t.Fatalf("备份码应可关闭,err=%v disabled=%v", err, store2.disabled) } // 启用 + 错误码 store3 := &fakeStore{enabled: true, secret: secret, backupOK: false} if err := newSvc(store3).Disable(context.Background(), "u1", "000000", "bad"); !errors.Is(err, ErrBadCode) { t.Fatalf("错误码应返回 ErrBadCode,实际: %v", err) } } func TestVerifyLogin(t *testing.T) { secret, _ := auth.GenerateTOTPSecret() code, _ := auth.TOTPCodeAt(secret, fixedNow) if !newSvc(&fakeStore{}).VerifyLogin(context.Background(), "u1", secret, code, "") { t.Fatal("正确 TOTP 应通过") } if !newSvc(&fakeStore{backupOK: true}).VerifyLogin(context.Background(), "u1", secret, "", "backup") { t.Fatal("有效备份码应通过") } if newSvc(&fakeStore{backupOK: false}).VerifyLogin(context.Background(), "u1", secret, "000000", "bad") { t.Fatal("错误码应不通过") } }