package mail import ( "context" "errors" "io" "log/slog" "testing" "gorm.io/gorm" "ws_usernode/internal/model" ) // flakyMailer 前 n 次发送失败,之后成功;用于验证重试。 type flakyMailer struct { failures int calls int } var errFlaky = errors.New("mail: smtp down") func (m *flakyMailer) Send(_ context.Context, _, _, _ string) error { m.calls++ if m.calls <= m.failures { return errFlaky } return nil } // recordingMailer 记录直发调用次数。 type recordingMailer struct{ calls int } func (m *recordingMailer) Send(context.Context, string, string, string) error { m.calls++ return nil } func testQueued(t *testing.T, base Mailer) (*QueuedMailer, *gorm.DB) { t.Helper() db, err := model.Open("sqlite", ":memory:", false) if err != nil { t.Fatalf("open test db: %v", err) } if err := db.AutoMigrate(&model.MailLog{}); err != nil { t.Fatalf("migrate mail_logs: %v", err) } log := slog.New(slog.NewTextHandler(io.Discard, nil)) return NewQueuedMailer(base, db, log), db } func TestQueuedMailerSendSuccess(t *testing.T) { rec := &recordingMailer{} q, db := testQueued(t, rec) if err := q.Send(context.Background(), "a@example.com", "主题", "正文"); err != nil { t.Fatalf("send: %v", err) } if rec.calls != 1 { t.Fatalf("base send calls = %d, want 1", rec.calls) } var n int64 if err := db.Model(&model.MailLog{}).Where("status = ?", StatusSent).Count(&n).Error; err != nil { t.Fatalf("count sent: %v", err) } if n != 1 { t.Fatalf("sent = %d, want 1", n) } } func TestQueuedMailerFailureQueuedAndRetry(t *testing.T) { base := &flakyMailer{failures: 2} q, db := testQueued(t, base) // 第一次发送失败:Send 仍返回 nil(不阻断业务),落 failed 记录 if err := q.Send(context.Background(), "b@example.com", "s", "body"); err != nil { t.Fatalf("send returned error: %v", err) } if base.calls != 1 { t.Fatalf("base calls = %d, want 1", base.calls) } var row model.MailLog if err := db.First(&row, "status = ?", StatusFailed).Error; err != nil { t.Fatalf("failed row: %v", err) } if row.RetryCount != 1 || row.Error != errFlaky.Error() { t.Fatalf("failed row = %+v", row) } // Retry 1:仍失败 → retry_count=2;Retry 2:成功 → sent if err := q.Retry(context.Background(), 5); err != nil { t.Fatalf("retry: %v", err) } if base.calls != 2 { t.Fatalf("calls after retry1 = %d, want 2", base.calls) } if err := q.Retry(context.Background(), 5); err != nil { t.Fatalf("retry2: %v", err) } if base.calls != 3 { t.Fatalf("calls after retry2 = %d, want 3", base.calls) } var n int64 if err := db.Model(&model.MailLog{}).Where("status = ?", StatusSent).Count(&n).Error; err != nil { t.Fatalf("count sent: %v", err) } if n != 1 { t.Fatalf("sent = %d, want 1", n) } } func TestQueuedMailerRetryHitsMax(t *testing.T) { base := &flakyMailer{failures: 100} // 永远失败 q, db := testQueued(t, base) if err := q.Send(context.Background(), "c@example.com", "s", "body"); err != nil { t.Fatalf("send: %v", err) } // 达到上限后 Retry 不再挑选(retry_count 不再增长) for i := 0; i < 10; i++ { if err := q.Retry(context.Background(), MaxRetries); err != nil { t.Fatalf("retry: %v", err) } } var row model.MailLog if err := db.First(&row, "status = ?", StatusFailed).Error; err != nil { t.Fatalf("failed row: %v", err) } if row.RetryCount > MaxRetries { t.Fatalf("retry_count = %d, want <= %d", row.RetryCount, MaxRetries) } }