first commit
This commit is contained in:
@@ -0,0 +1,212 @@
|
||||
// Package smtp implements the inbound SMTP MTA (port 25), submission
|
||||
// (port 587, STARTTLS + AUTH), and implicit-TLS SMTPS (port 465) — all as a
|
||||
// single hand-rolled state machine over net.Listener, per the project's
|
||||
// stdlib-first principle. No third-party SMTP library.
|
||||
package smtp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gomail/internal/config"
|
||||
"gomail/internal/db"
|
||||
"gomail/internal/mailstore"
|
||||
"gomail/internal/pipeline"
|
||||
"gomail/internal/ratelimit"
|
||||
)
|
||||
|
||||
const (
|
||||
maxCommandLine = 1000 // RFC 5321 command line limit
|
||||
maxRecipients = 100
|
||||
idleTimeout = 5 * time.Minute
|
||||
dataTimeout = 10 * time.Minute
|
||||
)
|
||||
|
||||
// Kind distinguishes the three listener roles — they share the same session
|
||||
// state machine but differ in whether TLS is implicit, STARTTLS-capable, or
|
||||
// plain (inbound MTA still offers STARTTLS, just doesn't require it for the
|
||||
// initial MAIL FROM the way submission does).
|
||||
type Kind int
|
||||
|
||||
const (
|
||||
KindMTA Kind = iota // :25 — inbound from the internet, STARTTLS optional
|
||||
KindSubmission // :587 — STARTTLS + AUTH required before MAIL FROM
|
||||
KindImplicitTLS // :465 — TLS from the first byte
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
cfg *config.Config
|
||||
database *db.DB
|
||||
store *mailstore.Store
|
||||
tlsConf *tls.Config
|
||||
pipeline *pipeline.Orchestrator // nil = pipeline disabled, all mail treated as clean
|
||||
|
||||
listeners []net.Listener
|
||||
wg sync.WaitGroup
|
||||
sessionWG sync.WaitGroup // tracks in-flight sessions for graceful drain
|
||||
|
||||
maxMessageBytes int64
|
||||
|
||||
connLimiter *ratelimit.Limiter // per-IP connections/min, cfg.RateLimits.SMTPConnPerMin
|
||||
authLimiter *ratelimit.Limiter // per-IP AUTH failures, cfg.RateLimits.SMTPAuthFailures (per minute)
|
||||
}
|
||||
|
||||
func NewServer(cfg *config.Config, database *db.DB, store *mailstore.Store, tlsConf *tls.Config, orch *pipeline.Orchestrator) *Server {
|
||||
return &Server{
|
||||
cfg: cfg,
|
||||
database: database,
|
||||
store: store,
|
||||
tlsConf: tlsConf,
|
||||
pipeline: orch,
|
||||
maxMessageBytes: int64(cfg.Storage.MaxMessageSizeMB) * 1024 * 1024,
|
||||
connLimiter: ratelimit.New(cfg.RateLimits.SMTPConnPerMin),
|
||||
authLimiter: ratelimit.New(cfg.RateLimits.SMTPAuthFailures),
|
||||
}
|
||||
}
|
||||
|
||||
// ListenAndServe starts all three listeners and blocks until one fails or
|
||||
// ctx is cancelled. Each listener's accept loop runs in its own goroutine.
|
||||
func (s *Server) ListenAndServe(ctx context.Context) error {
|
||||
specs := []struct {
|
||||
addr string
|
||||
kind Kind
|
||||
}{
|
||||
{s.cfg.Server.SMTPAddr, KindMTA},
|
||||
{s.cfg.Server.SubmissionAddr, KindSubmission},
|
||||
{s.cfg.Server.SMTPSAddr, KindImplicitTLS},
|
||||
}
|
||||
|
||||
errCh := make(chan error, len(specs))
|
||||
|
||||
for _, spec := range specs {
|
||||
ln, err := net.Listen("tcp", spec.addr)
|
||||
if err != nil {
|
||||
s.closeAll()
|
||||
return fmt.Errorf("listen %s: %w", spec.addr, err)
|
||||
}
|
||||
|
||||
if spec.kind == KindImplicitTLS {
|
||||
ln = tls.NewListener(ln, s.tlsConf)
|
||||
}
|
||||
|
||||
s.listeners = append(s.listeners, ln)
|
||||
slog.Info("SMTP listener started", "addr", spec.addr, "kind", kindName(spec.kind))
|
||||
|
||||
s.wg.Add(1)
|
||||
go func(ln net.Listener, kind Kind) {
|
||||
defer s.wg.Done()
|
||||
s.acceptLoop(ctx, ln, kind)
|
||||
}(ln, spec.kind)
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case err := <-errCh:
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) acceptLoop(ctx context.Context, ln net.Listener, kind Kind) {
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return // expected — listener closed during shutdown
|
||||
default:
|
||||
slog.Error("accept error", "err", err, "kind", kindName(kind))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
ip := connHost(conn.RemoteAddr())
|
||||
if !s.connLimiter.Allow(ip) {
|
||||
slog.Warn("SMTP connection rate limit exceeded, rejecting", "ip", ip, "kind", kindName(kind))
|
||||
conn.Close()
|
||||
continue
|
||||
}
|
||||
|
||||
s.sessionWG.Add(1)
|
||||
go func() {
|
||||
defer s.sessionWG.Done()
|
||||
s.handleConn(ctx, conn, kind)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleConn(ctx context.Context, conn net.Conn, kind Kind) {
|
||||
defer conn.Close()
|
||||
|
||||
sess := &session{
|
||||
conn: conn,
|
||||
server: s,
|
||||
kind: kind,
|
||||
hostname: s.cfg.Server.Hostname,
|
||||
}
|
||||
|
||||
remoteAddr := conn.RemoteAddr()
|
||||
if tcpAddr, ok := remoteAddr.(*net.TCPAddr); ok {
|
||||
sess.senderIP = tcpAddr.IP
|
||||
}
|
||||
|
||||
slog.Debug("SMTP connection accepted", "remote", remoteAddr, "kind", kindName(kind))
|
||||
sess.run(ctx)
|
||||
}
|
||||
|
||||
// Shutdown closes all listeners immediately (stops accepting new
|
||||
// connections) then waits up to gracePeriod for in-flight sessions to finish
|
||||
// naturally (they'll see ctx.Done() and wind down at their next command read).
|
||||
func (s *Server) Shutdown(gracePeriod time.Duration) {
|
||||
s.closeAll()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
s.sessionWG.Wait()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
slog.Info("all SMTP sessions drained cleanly")
|
||||
case <-time.After(gracePeriod):
|
||||
slog.Warn("SMTP shutdown grace period expired — some sessions forcibly terminated", "grace_period", gracePeriod)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) closeAll() {
|
||||
for _, ln := range s.listeners {
|
||||
ln.Close()
|
||||
}
|
||||
s.wg.Wait()
|
||||
}
|
||||
|
||||
func kindName(k Kind) string {
|
||||
switch k {
|
||||
case KindMTA:
|
||||
return "mta"
|
||||
case KindSubmission:
|
||||
return "submission"
|
||||
case KindImplicitTLS:
|
||||
return "smtps"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
// connHost extracts just the IP (no port) from a net.Addr, for use as a
|
||||
// rate-limiter key — falls back to the full address string if it isn't
|
||||
// host:port shaped (shouldn't happen for real TCP connections, but a
|
||||
// fallback beats a panic).
|
||||
func connHost(addr net.Addr) string {
|
||||
host, _, err := net.SplitHostPort(addr.String())
|
||||
if err != nil {
|
||||
return addr.String()
|
||||
}
|
||||
return host
|
||||
}
|
||||
Reference in New Issue
Block a user