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