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") messageID := cleanHeader(message.ID) + "@aigw.local" _, writeErr := fmt.Fprintf(w, "From: %s\r\nTo: %s\r\nSubject: %s\r\nMessage-ID: <%s>\r\nX-AIGW-Message-ID: %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), messageID, cleanHeader(message.ID), 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) } } }