summaryrefslogtreecommitdiffstats
path: root/internal/smtpclient/client.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/smtpclient/client.go')
-rw-r--r--internal/smtpclient/client.go208
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
+}