diff options
Diffstat (limited to 'internal/mailer')
| -rw-r--r-- | internal/mailer/mailer.go | 170 |
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) + } + } +} |
