package postcard import ( "context" "fmt" "log/slog" netmail "net/mail" "strings" "sync" "time" gomail "github.com/wneessen/go-mail" ) // Driver delivers a fully rendered message. type Driver interface { Send(ctx context.Context, msg RenderedMessage) error } // RenderedMessage is the validated payload a driver transmits. type RenderedMessage struct { To []string Cc []string Bcc []string From string ReplyTo string Subject string HTML string Text string } // SMTPConfig is the compass mail.smtp section used by the SMTP driver. type SMTPConfig struct { Host string Port int Username string Password string TLS string Timeout time.Duration } // MemoryDriver stores rendered messages for tests. It is concurrency-safe. type MemoryDriver struct { mu sync.Mutex messages []RenderedMessage } // NewMemoryDriver returns an empty in-memory driver. func NewMemoryDriver() *MemoryDriver { return &MemoryDriver{} } // Send appends msg to the in-memory log. func (d *MemoryDriver) Send(_ context.Context, msg RenderedMessage) error { if d == nil { return fmt.Errorf("postcard: memory driver is nil") } d.mu.Lock() defer d.mu.Unlock() d.messages = append(d.messages, cloneRendered(msg)) return nil } // Messages returns a copy of stored messages in send order. func (d *MemoryDriver) Messages() []RenderedMessage { if d == nil { return nil } d.mu.Lock() defer d.mu.Unlock() out := make([]RenderedMessage, len(d.messages)) for i, msg := range d.messages { out[i] = cloneRendered(msg) } return out } // LogDriver writes rendered headers and the text part to a logger. type LogDriver struct { log *slog.Logger } // NewLogDriver logs through log. slog.Default is used when log is nil. func NewLogDriver(log *slog.Logger) *LogDriver { if log == nil { log = slog.Default() } return &LogDriver{log: log} } // Send logs From/To/Cc/Bcc/ReplyTo/Subject and the text part. It never // logs SMTP credentials or the HTML body. func (d *LogDriver) Send(_ context.Context, msg RenderedMessage) error { if d == nil || d.log == nil { return fmt.Errorf("postcard: log driver is nil") } d.log.Info("postcard mail", slog.String("from", msg.From), slog.Any("to", msg.To), slog.Any("cc", msg.Cc), slog.Any("bcc", msg.Bcc), slog.String("reply_to", msg.ReplyTo), slog.String("subject", msg.Subject), slog.String("text", msg.Text), ) return nil } // SMTPDriver sends through go-mail with an explicit TLS policy. type SMTPDriver struct { client *gomail.Client } // NewSMTPDriver builds a go-mail client. TLS defaults to mandatory; NoTLS // is opt-in for local Mailpit and is never inferred. func NewSMTPDriver(cfg SMTPConfig) (*SMTPDriver, error) { if strings.TrimSpace(cfg.Host) == "" { return nil, fmt.Errorf("postcard: mail.smtp.host is empty") } port := cfg.Port if port == 0 { port = 587 } timeout := cfg.Timeout if timeout <= 0 { timeout = 10 * time.Second } policy, err := parseTLSPolicy(cfg.TLS) if err != nil { return nil, err } opts := []gomail.Option{ gomail.WithPort(port), gomail.WithTLSPolicy(policy), gomail.WithTimeout(timeout), } if strings.TrimSpace(cfg.Username) != "" { opts = append(opts, gomail.WithSMTPAuth(gomail.SMTPAuthPlain), gomail.WithUsername(cfg.Username), gomail.WithPassword(cfg.Password), ) } client, err := gomail.NewClient(cfg.Host, opts...) if err != nil { return nil, fmt.Errorf("postcard: smtp client: %w", err) } return &SMTPDriver{client: client}, nil } // Send validates headers and recipients, then delivers with context. func (d *SMTPDriver) Send(ctx context.Context, msg RenderedMessage) error { if d == nil || d.client == nil { return fmt.Errorf("postcard: smtp driver is nil") } if err := validateRendered(msg, true); err != nil { return err } gm, err := buildSMTPMessage(msg) if err != nil { return err } if err := d.client.DialAndSendWithContext(ctx, gm); err != nil { return fmt.Errorf("smtp: %w", err) } return nil } // FailDriver returns Err from Send. Tests use it as a deterministic failure. type FailDriver struct { Err error sends int sendsMu sync.Mutex } // Send returns Err without retrying. func (d *FailDriver) Send(context.Context, RenderedMessage) error { if d == nil { return fmt.Errorf("postcard: fail driver is nil") } d.sendsMu.Lock() d.sends++ d.sendsMu.Unlock() if d.Err == nil { return fmt.Errorf("postcard: driver failed") } return d.Err } func (d *FailDriver) sendCount() int { d.sendsMu.Lock() defer d.sendsMu.Unlock() return d.sends } func cloneRendered(msg RenderedMessage) RenderedMessage { msg.To = append([]string(nil), msg.To...) msg.Cc = append([]string(nil), msg.Cc...) msg.Bcc = append([]string(nil), msg.Bcc...) return msg } func parseTLSPolicy(v string) (gomail.TLSPolicy, error) { switch strings.ToLower(strings.TrimSpace(v)) { case "", "mandatory", "tls": return gomail.TLSMandatory, nil case "opportunistic", "starttls": return gomail.TLSOpportunistic, nil case "none", "notls", "off": return gomail.NoTLS, nil default: return 0, fmt.Errorf("postcard: unknown mail.smtp.tls %q (use mandatory, starttls, or none)", v) } } func buildSMTPMessage(msg RenderedMessage) (*gomail.Msg, error) { gm := gomail.NewMsg() if err := gm.From(msg.From); err != nil { return nil, fmt.Errorf("postcard: from: %w", err) } if len(msg.To) > 0 { if err := gm.To(msg.To...); err != nil { return nil, fmt.Errorf("postcard: to: %w", err) } } if len(msg.Cc) > 0 { if err := gm.Cc(msg.Cc...); err != nil { return nil, fmt.Errorf("postcard: cc: %w", err) } } if len(msg.Bcc) > 0 { if err := gm.Bcc(msg.Bcc...); err != nil { return nil, fmt.Errorf("postcard: bcc: %w", err) } } if msg.ReplyTo != "" { if err := gm.ReplyTo(msg.ReplyTo); err != nil { return nil, fmt.Errorf("postcard: reply-to: %w", err) } } gm.Subject(msg.Subject) gm.SetBodyString(gomail.TypeTextPlain, msg.Text) gm.AddAlternativeString(gomail.TypeTextHTML, msg.HTML) return gm, nil } func validateRendered(msg RenderedMessage, requireFrom bool) error { if err := rejectCRLF("subject", msg.Subject); err != nil { return err } if err := rejectCRLF("from", msg.From); err != nil { return err } if requireFrom || msg.From != "" { if err := validateAddr("from", msg.From); err != nil { return err } } if err := rejectCRLF("reply-to", msg.ReplyTo); err != nil { return err } if msg.ReplyTo != "" { if err := validateAddr("reply-to", msg.ReplyTo); err != nil { return err } } for _, addr := range msg.To { if err := validateAddr("to", addr); err != nil { return err } } for _, addr := range msg.Cc { if err := validateAddr("cc", addr); err != nil { return err } } for _, addr := range msg.Bcc { if err := validateAddr("bcc", addr); err != nil { return err } } return nil } func validateAddr(kind, addr string) error { if err := rejectCRLF(kind, addr); err != nil { return err } if strings.TrimSpace(addr) == "" { return fmt.Errorf("postcard: empty %s address", kind) } if _, err := netmail.ParseAddress(addr); err != nil { return fmt.Errorf("postcard: invalid %s address", kind) } return nil } func rejectCRLF(kind, value string) error { if strings.ContainsAny(value, "\r\n") { return fmt.Errorf("postcard: %s contains CR/LF", kind) } return nil }