package inbound import ( "bytes" "context" "crypto/tls" "database/sql" "errors" "fmt" "io" "log/slog" "net" "net/mail" "os" "strings" "sync" "time" "postern/internal/db" "postern/internal/model" "postern/internal/policy" "github.com/emersion/go-imap/v2/imapserver" "github.com/emersion/go-sasl" smtpserver "github.com/emersion/go-smtp" "github.com/google/uuid" ) type MessageDeliverer interface { SubmitMessage(ctx context.Context, conn net.Conn, user model.User, r io.Reader) error } type Config struct { DB *db.DB Deliverer MessageDeliverer Listen string TLS *tls.Config PolicyEngine *policy.Engine } type Server struct { db *db.DB listen string deliverer MessageDeliverer tlsConfig *tls.Config policyEngine *policy.Engine mu sync.Mutex smtpSrv *smtpserver.Server status string } func NewServer(config *Config) (*Server, error) { return &Server{ db: config.DB, listen: config.Listen, deliverer: config.Deliverer, tlsConfig: config.TLS, policyEngine: config.PolicyEngine, }, nil } func (i *Server) Start() error { be := &backend{ db: i.db, deliverer: i.deliverer, policyEngine: i.policyEngine, } s := smtpserver.NewServer(be) s.Addr = i.listen s.AllowInsecureAuth = false s.TLSConfig = i.tlsConfig s.MaxMessageBytes = 20_000_000 i.mu.Lock() i.smtpSrv = s i.status = "running" i.mu.Unlock() err := s.ListenAndServe() if err != nil { return err } return nil } func (i *Server) Stop() error { i.mu.Lock() defer i.mu.Unlock() i.status = "stopped" if i.smtpSrv != nil { return i.smtpSrv.Close() } return nil } func (i *Server) Status() string { i.mu.Lock() defer i.mu.Unlock() return i.status } type smtpSession struct { db *db.DB conn *smtpserver.Conn deliverer MessageDeliverer authenticatedUser *model.SMTPAuth toAddress string fromAddress string helo string destinationUser model.User sessionID string policyEngine *policy.Engine logger *slog.Logger } // The backend implements SMTP server methods. type backend struct { db *db.DB deliverer MessageDeliverer policyEngine *policy.Engine } // NewSession is called after client greeting (EHLO, HELO). func (bkd *backend) NewSession(c *smtpserver.Conn) (smtpserver.Session, error) { sessionID := uuid.NewString() l := slog.With("component", "inbound", "remote_ip", c.Conn().RemoteAddr().String(), "session_id", sessionID) l.Info("new session") return &smtpSession{ db: bkd.db, conn: c, deliverer: bkd.deliverer, authenticatedUser: &model.SMTPAuth{}, sessionID: sessionID, logger: l, policyEngine: bkd.policyEngine, }, nil } // AuthMechanisms returns a slice of available auth mechanisms func (s *smtpSession) AuthMechanisms() []string { return []string{sasl.Plain} } // Auth is the handler for supported authenticators. func (s *smtpSession) Auth(_ string) (sasl.Server, error) { return sasl.NewPlainServer(func(identity, username, password string) error { u, err := s.db.GetSMTPAuthUser(context.Background(), username) if err != nil { if errors.Is(err, sql.ErrNoRows) { s.logger.Info("user not found") } else { s.logger.Error("getting user", "error", err) } return imapserver.ErrAuthFailed } err = u.VerifyPassword([]byte(password)) if err != nil { slog.Info("invalid password", "username", username) return imapserver.ErrAuthFailed } s.authenticatedUser = &u s.logger.Info("authenticated", "user_id", u.ID, "username", username) return nil }), nil } func (s *smtpSession) Reset() { s.toAddress = "" s.fromAddress = "" } func (s *smtpSession) Logout() error { s.authenticatedUser = nil return nil } func (s *smtpSession) Mail(from string, _ *smtpserver.MailOptions) error { if s.authenticatedUser == nil { return smtpserver.ErrAuthFailed } s.logger.Info("MAIL FROM", "from", from) s.fromAddress = from return nil } func (s *smtpSession) Rcpt(to string, _ *smtpserver.RcptOptions) error { if s.authenticatedUser == nil { return smtpserver.ErrAuthFailed } ctx := context.Background() s.logger.InfoContext(ctx, "RCPT TO", "to", to) addr, err := mail.ParseAddress(to) if err != nil { s.logger.InfoContext(ctx, "RCPT TO", "error", err) return &smtpserver.SMTPError{ Code: 501, EnhancedCode: smtpserver.EnhancedCode{5, 1, 3}, Message: "Bad recipient address syntax", } } atIndex := strings.LastIndex(addr.Address, "@") if atIndex == -1 { s.logger.InfoContext(ctx, "RCPT TO", "error", "malformed email: missing @ domain separator") return &smtpserver.SMTPError{ Code: 501, EnhancedCode: smtpserver.EnhancedCode{5, 1, 3}, Message: "Bad recipient address syntax", } } localPart := addr.Address[:atIndex] domain := addr.Address[atIndex+1:] baseName, tag, _ := strings.Cut(localPart, "+") s.logger.InfoContext(ctx, "parsed address", "baseName", baseName, "tag", tag, "domain", domain) u, err := s.db.GetUserForAddress(ctx, fmt.Sprintf("%s@%s", baseName, domain)) if err != nil { if errors.Is(err, sql.ErrNoRows) { s.logger.InfoContext(ctx, "RCPT TO address not found") return &smtpserver.SMTPError{ Code: 550, EnhancedCode: smtpserver.EnhancedCode{5, 1, 1}, Message: "Mailbox unavailable", } } s.logger.ErrorContext(ctx, "RCPT TO get address", "error", err) return &smtpserver.SMTPError{ Code: 451, EnhancedCode: smtpserver.EnhancedCode{4, 3, 0}, Message: "Temporary local problem, please try again later", } } s.destinationUser = u s.toAddress = to s.logger.InfoContext(ctx, "RCPT TO", "destination_user", u.Name, "destination_user_id", u.ID) return nil } func (s *smtpSession) Data(r io.Reader) error { if s.authenticatedUser == nil { return smtpserver.ErrAuthFailed } ctx := context.Background() s.logger.InfoContext(ctx, "DATA") var buffer bytes.Buffer // Prepend headers buffer.WriteString(fmt.Sprintf("Return-Path: %s\r\n", s.fromAddress)) myHostname, err := os.Hostname() if err != nil { myHostname = "postern.local" } with := "ESMTPSA" // we only allow encrypted (S) and authenticated (A) connections for inbound receivedTimestampLayout := "Mon, 02 Jan 2006 15:04:05 -0700 (UTC)" receivedTimestamp := time.Now().UTC().Format(receivedTimestampLayout) buffer.WriteString(fmt.Sprintf("Received: from %s (%s) by %s (Postern p25.dev) with %s id %s for <%s>; %s\r\n", s.conn.Hostname(), s.conn.Conn().RemoteAddr().String(), myHostname, with, s.sessionID, s.toAddress, receivedTimestamp)) _, err = io.Copy(&buffer, r) if err != nil { s.logger.ErrorContext(ctx, "Copy", "error", err) return errInternalServerError } message := buffer.Bytes() err = s.deliverer.SubmitMessage(ctx, s.conn.Conn(), s.destinationUser, bytes.NewReader(message)) if err != nil { if errors.Is(err, model.ErrMissingFromHeader) { return &smtpserver.SMTPError{ Code: 550, EnhancedCode: smtpserver.EnhancedCode{5, 7, 1}, Message: "From header is required but missing", } } s.logger.ErrorContext(ctx, "SubmitMessage", "error", err) return errInternalServerError } return nil }