diff options
Diffstat (limited to 'internal/smtpclient/client.go')
| -rw-r--r-- | internal/smtpclient/client.go | 208 |
1 files changed, 208 insertions, 0 deletions
diff --git a/internal/smtpclient/client.go b/internal/smtpclient/client.go new file mode 100644 index 0000000..d4ee686 --- /dev/null +++ b/internal/smtpclient/client.go @@ -0,0 +1,208 @@ +package smtpclient + +import ( + "bufio" + "context" + "crypto/tls" + "fmt" + "io" + "net" + "net/textproto" + "strings" + "time" +) + +type DialContextFunc func(ctx context.Context, network, address string) (net.Conn, error) + +type Config struct { + Host string + Port int + Recipient string + EnvelopeFrom string + HELO string + TLSServerName string + RequireTLS bool + ImplicitTLS bool + Timeout time.Duration + DryRun bool +} + +type Client struct { + cfg Config + dial DialContextFunc +} + +type Message struct { + EnvelopeFrom string + Raw string +} + +func New(cfg Config, dial DialContextFunc) *Client { + return &Client{cfg: cfg, dial: dial} +} + +func (c *Client) Send(ctx context.Context, msg Message) error { + if c.cfg.DryRun { + return nil + } + if c.dial == nil { + c.dial = (&net.Dialer{}).DialContext + } + timeout := c.cfg.Timeout + if timeout <= 0 { + timeout = 90 * time.Second + } + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + + addr := net.JoinHostPort(c.cfg.Host, fmt.Sprintf("%d", c.cfg.Port)) + conn, err := c.dial(ctx, "tcp", addr) + if err != nil { + return fmt.Errorf("dial smtp: %w", err) + } + defer conn.Close() + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetDeadline(deadline) + } + if c.cfg.ImplicitTLS { + tlsConn := tls.Client(conn, &tls.Config{ + ServerName: c.cfg.TLSServerName, + MinVersion: tls.VersionTLS12, + }) + if err := tlsConn.HandshakeContext(ctx); err != nil { + return fmt.Errorf("implicit tls handshake: %w", err) + } + conn = tlsConn + } + + session := newSession(conn) + if _, _, err := session.read(220); err != nil { + return fmt.Errorf("smtp greeting: %w", err) + } + + if err := session.ehlo(c.cfg.HELO); err != nil { + return err + } + + if c.cfg.RequireTLS && !c.cfg.ImplicitTLS { + if _, _, err := session.cmd(220, "STARTTLS\r\n"); err != nil { + return fmt.Errorf("starttls: %w", err) + } + tlsConn := tls.Client(conn, &tls.Config{ + ServerName: c.cfg.TLSServerName, + MinVersion: tls.VersionTLS12, + }) + if err := tlsConn.HandshakeContext(ctx); err != nil { + return fmt.Errorf("tls handshake: %w", err) + } + session = newSession(tlsConn) + if err := session.ehlo(c.cfg.HELO); err != nil { + return err + } + } + + envelopeFrom := sanitizeEnvelope(c.cfg.EnvelopeFrom) + if envelopeFrom == "" { + envelopeFrom = sanitizeEnvelope(msg.EnvelopeFrom) + } + if _, _, err := session.cmd(250, "MAIL FROM:<%s>\r\n", envelopeFrom); err != nil { + return fmt.Errorf("mail from rejected: %w", err) + } + if _, _, err := session.cmd(250, "RCPT TO:<%s>\r\n", sanitizeEnvelope(c.cfg.Recipient)); err != nil { + return fmt.Errorf("rcpt to rejected: %w", err) + } + if _, _, err := session.cmd(354, "DATA\r\n"); err != nil { + return fmt.Errorf("data rejected: %w", err) + } + if err := writeSMTPData(session.w, msg.Raw); err != nil { + return err + } + if _, _, err := session.read(250); err != nil { + return fmt.Errorf("message rejected: %w", err) + } + _, _, _ = session.cmd(221, "QUIT\r\n") + return nil +} + +type session struct { + conn net.Conn + tp *textproto.Reader + w *bufio.Writer +} + +func newSession(conn net.Conn) *session { + reader := bufio.NewReader(conn) + return &session{ + conn: conn, + tp: textproto.NewReader(reader), + w: bufio.NewWriter(conn), + } +} + +func (s *session) ehlo(helo string) error { + if helo == "" { + helo = "n2usenet.local" + } + if _, _, err := s.cmd(250, "EHLO %s\r\n", sanitizeAtom(helo)); err != nil { + if _, _, heloErr := s.cmd(250, "HELO %s\r\n", sanitizeAtom(helo)); heloErr != nil { + return fmt.Errorf("ehlo failed: %w", err) + } + } + return nil +} + +func (s *session) cmd(expect int, format string, args ...any) (int, string, error) { + if _, err := fmt.Fprintf(s.w, format, args...); err != nil { + return 0, "", err + } + if err := s.w.Flush(); err != nil { + return 0, "", err + } + return s.read(expect) +} + +func (s *session) read(expect int) (int, string, error) { + code, msg, err := s.tp.ReadResponse(expect) + if err != nil { + return code, msg, err + } + return code, msg, nil +} + +func writeSMTPData(w *bufio.Writer, raw string) error { + raw = strings.ReplaceAll(raw, "\r\n", "\n") + raw = strings.ReplaceAll(raw, "\r", "\n") + for _, line := range strings.Split(raw, "\n") { + if strings.HasPrefix(line, ".") { + line = "." + line + } + if _, err := io.WriteString(w, line+"\r\n"); err != nil { + return fmt.Errorf("write smtp data: %w", err) + } + } + if _, err := io.WriteString(w, ".\r\n"); err != nil { + return fmt.Errorf("write smtp terminator: %w", err) + } + if err := w.Flush(); err != nil { + return fmt.Errorf("flush smtp data: %w", err) + } + return nil +} + +func sanitizeEnvelope(v string) string { + v = strings.TrimSpace(v) + v = strings.ReplaceAll(v, "\r", "") + v = strings.ReplaceAll(v, "\n", "") + v = strings.Trim(v, "<>") + return v +} + +func sanitizeAtom(v string) string { + v = strings.TrimSpace(v) + v = strings.ReplaceAll(v, "\r", "") + v = strings.ReplaceAll(v, "\n", "") + if v == "" { + return "n2usenet.local" + } + return v +} |
