summaryrefslogtreecommitdiff
path: root/internal/mailer/mailer.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/mailer/mailer.go170
1 files changed, 170 insertions, 0 deletions
diff --git a/internal/mailer/mailer.go b/internal/mailer/mailer.go
new file mode 100644
index 0000000..4ec11ac
--- /dev/null
+++ b/internal/mailer/mailer.go
@@ -0,0 +1,170 @@
+package mailer
+
+import (
+ "context"
+ "crypto/tls"
+ "errors"
+ "fmt"
+ "log/slog"
+ "net"
+ "net/smtp"
+ "strings"
+ "time"
+)
+
+type Message struct {
+ ID string
+ Recipient string
+ Subject string
+ Body string
+}
+
+type Queue interface {
+ ClaimMail(context.Context) (Message, bool, error)
+ MarkMailSent(context.Context, string) error
+ MarkMailFailed(context.Context, string, error) error
+}
+
+type SMTPConfig struct {
+ Address string
+ Username string
+ Password string
+ FromName string
+ FromAddress string
+ TLSMode string
+ DialTimeout time.Duration
+}
+
+type Sender struct {
+ config SMTPConfig
+}
+
+func NewSender(config SMTPConfig) (*Sender, error) {
+ if strings.TrimSpace(config.Address) == "" || strings.TrimSpace(config.FromAddress) == "" {
+ return nil, errors.New("SMTP address and from address are required")
+ }
+ if config.DialTimeout <= 0 {
+ config.DialTimeout = 10 * time.Second
+ }
+ return &Sender{config: config}, nil
+}
+
+func (s *Sender) Send(ctx context.Context, message Message) error {
+ host, _, err := net.SplitHostPort(s.config.Address)
+ if err != nil {
+ return fmt.Errorf("parse SMTP address: %w", err)
+ }
+ dialer := net.Dialer{Timeout: s.config.DialTimeout}
+ var conn net.Conn
+ if s.config.TLSMode == "tls" {
+ conn, err = tls.DialWithDialer(&dialer, "tcp", s.config.Address, &tls.Config{MinVersion: tls.VersionTLS12, ServerName: host})
+ } else {
+ conn, err = dialer.DialContext(ctx, "tcp", s.config.Address)
+ }
+ if err != nil {
+ return fmt.Errorf("connect SMTP: %w", err)
+ }
+ defer conn.Close()
+ client, err := smtp.NewClient(conn, host)
+ if err != nil {
+ return fmt.Errorf("open SMTP client: %w", err)
+ }
+ defer client.Close()
+ if s.config.TLSMode == "starttls" {
+ if ok, _ := client.Extension("STARTTLS"); !ok {
+ return errors.New("SMTP server does not support required STARTTLS")
+ }
+ if err := client.StartTLS(&tls.Config{MinVersion: tls.VersionTLS12, ServerName: host}); err != nil {
+ return fmt.Errorf("start SMTP TLS: %w", err)
+ }
+ }
+ if s.config.Username != "" {
+ if ok, _ := client.Extension("AUTH"); !ok {
+ return errors.New("SMTP server does not support authentication")
+ }
+ if err := client.Auth(smtp.PlainAuth("", s.config.Username, s.config.Password, host)); err != nil {
+ return fmt.Errorf("authenticate SMTP: %w", err)
+ }
+ }
+ if err := client.Mail(s.config.FromAddress); err != nil {
+ return fmt.Errorf("set SMTP sender: %w", err)
+ }
+ if err := client.Rcpt(message.Recipient); err != nil {
+ return fmt.Errorf("set SMTP recipient: %w", err)
+ }
+ w, err := client.Data()
+ if err != nil {
+ return fmt.Errorf("start SMTP body: %w", err)
+ }
+ from := s.config.FromAddress
+ if name := cleanHeader(s.config.FromName); name != "" {
+ from = fmt.Sprintf("%s <%s>", name, s.config.FromAddress)
+ }
+ body := strings.ReplaceAll(message.Body, "\r\n", "\n")
+ body = strings.ReplaceAll(body, "\n", "\r\n")
+ _, writeErr := fmt.Fprintf(w, "From: %s\r\nTo: %s\r\nSubject: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: 8bit\r\n\r\n%s\r\n",
+ from, cleanHeader(message.Recipient), cleanHeader(message.Subject), body)
+ closeErr := w.Close()
+ if writeErr != nil {
+ return fmt.Errorf("write SMTP body: %w", writeErr)
+ }
+ if closeErr != nil {
+ return fmt.Errorf("finish SMTP body: %w", closeErr)
+ }
+ if err := client.Quit(); err != nil {
+ return fmt.Errorf("finish SMTP session: %w", err)
+ }
+ return nil
+}
+
+func cleanHeader(value string) string {
+ value = strings.ReplaceAll(value, "\r", "")
+ return strings.ReplaceAll(value, "\n", "")
+}
+
+type Worker struct {
+ queue Queue
+ sender *Sender
+ logger *slog.Logger
+}
+
+func NewWorker(queue Queue, sender *Sender, logger *slog.Logger) *Worker {
+ if logger == nil {
+ logger = slog.Default()
+ }
+ return &Worker{queue: queue, sender: sender, logger: logger}
+}
+
+func (w *Worker) Run(ctx context.Context) {
+ ticker := time.NewTicker(3 * time.Second)
+ defer ticker.Stop()
+ for {
+ w.drain(ctx)
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ }
+ }
+}
+
+func (w *Worker) drain(ctx context.Context) {
+ for i := 0; i < 20 && ctx.Err() == nil; i++ {
+ message, ok, err := w.queue.ClaimMail(ctx)
+ if err != nil {
+ w.logger.Warn("mail_outbox_claim_failed", "error", err)
+ return
+ }
+ if !ok {
+ return
+ }
+ if err := w.sender.Send(ctx, message); err != nil {
+ _ = w.queue.MarkMailFailed(ctx, message.ID, err)
+ w.logger.Warn("mail_delivery_failed", "message_id", message.ID, "error", err)
+ continue
+ }
+ if err := w.queue.MarkMailSent(ctx, message.ID); err != nil {
+ w.logger.Warn("mail_outbox_complete_failed", "message_id", message.ID, "error", err)
+ }
+ }
+}