fix(email): unify SMTP connection path between send and test-connection

- UseTLS now tries implicit TLS first (port 465 semantics) and, when the
  server answers in plaintext (tls.RecordHeaderError, e.g. port 587
  submission), automatically retries with mandatory STARTTLS; encryption
  is never silently downgraded (fixes #1470, supersedes #1488)
- TestSMTPConnectionWithConfig now shares connectSMTP with the send path,
  adding the opportunistic STARTTLS upgrade the send path gained in
  b402c367d; this removes the 'test connection fails but test email
  sends' mismatch reported in #1488
- test-connection now also honors dial/IO timeouts and ignores
  non-standard QUIT responses, matching the send path
This commit is contained in:
shaw
2026-07-31 23:12:02 +08:00
parent 570ea74d12
commit 4c80d160dd
2 changed files with 475 additions and 119 deletions
+92 -119
View File
@@ -5,7 +5,9 @@ import (
"crypto/rand"
"crypto/subtle"
"crypto/tls"
"crypto/x509"
"encoding/hex"
"errors"
"fmt"
"html"
"log/slog"
@@ -191,118 +193,114 @@ func (s *EmailService) SendEmailWithConfig(config *SMTPConfig, to, subject, body
return err
}
addr := fmt.Sprintf("%s:%d", config.Host, config.Port)
auth := smtp.PlainAuth("", config.Username, config.Password, config.Host)
if config.UseTLS {
return s.sendMailTLS(addr, auth, message.envelopeFrom, message.envelopeTo, message.data, config.Host)
}
return s.sendMailPlain(addr, auth, message.envelopeFrom, message.envelopeTo, message.data, config.Host)
}
// sendMailPlain sends mail without TLS using a dialer with timeout.
func (s *EmailService) sendMailPlain(addr string, auth smtp.Auth, from, to string, msg []byte, host string) error {
dialer := &net.Dialer{Timeout: smtpDialTimeout}
conn, err := dialer.Dial("tcp", addr)
client, err := s.connectSMTP(config)
if err != nil {
return fmt.Errorf("smtp dial: %w", err)
}
_ = conn.SetDeadline(time.Now().Add(smtpIOTimeout))
defer func() { _ = conn.Close() }()
client, err := smtp.NewClient(conn, host)
if err != nil {
return fmt.Errorf("new smtp client: %w", err)
return err
}
defer func() { _ = client.Close() }()
// Opportunistic STARTTLS: upgrade to encrypted connection if the server supports it.
// This mirrors the behavior of smtp.SendMail which we replaced for timeout support.
if ok, _ := client.Extension("STARTTLS"); ok {
if err = client.StartTLS(&tls.Config{ServerName: host, MinVersion: tls.VersionTLS12}); err != nil {
return fmt.Errorf("starttls: %w", err)
}
}
auth := smtp.PlainAuth("", config.Username, config.Password, config.Host)
if err = client.Auth(auth); err != nil {
return fmt.Errorf("smtp auth: %w", err)
}
if err = client.Mail(from); err != nil {
if err = client.Mail(message.envelopeFrom); err != nil {
return fmt.Errorf("smtp mail: %w", err)
}
if err = client.Rcpt(to); err != nil {
if err = client.Rcpt(message.envelopeTo); err != nil {
return fmt.Errorf("smtp rcpt: %w", err)
}
w, err := client.Data()
if err != nil {
return fmt.Errorf("smtp data: %w", err)
}
if _, err = w.Write(msg); err != nil {
if _, err = w.Write(message.data); err != nil {
return fmt.Errorf("write msg: %w", err)
}
if err = w.Close(); err != nil {
return fmt.Errorf("close writer: %w", err)
}
_ = client.Quit()
return nil
}
// sendMailTLS 使用TLS发送邮件
func (s *EmailService) sendMailTLS(addr string, auth smtp.Auth, from, to string, msg []byte, host string) error {
tlsConfig := &tls.Config{
ServerName: host,
// 强制 TLS 1.2+,避免协议降级导致的弱加密风险。
MinVersion: tls.VersionTLS12,
}
dialer := &net.Dialer{Timeout: smtpDialTimeout}
conn, err := tls.DialWithDialer(dialer, "tcp", addr, tlsConfig)
if err != nil {
return fmt.Errorf("tls dial: %w", err)
}
_ = conn.SetDeadline(time.Now().Add(smtpIOTimeout))
defer func() { _ = conn.Close() }()
client, err := smtp.NewClient(conn, host)
if err != nil {
return fmt.Errorf("new smtp client: %w", err)
}
defer func() { _ = client.Close() }()
if err = client.Auth(auth); err != nil {
return fmt.Errorf("smtp auth: %w", err)
}
if err = client.Mail(from); err != nil {
return fmt.Errorf("smtp mail: %w", err)
}
if err = client.Rcpt(to); err != nil {
return fmt.Errorf("smtp rcpt: %w", err)
}
w, err := client.Data()
if err != nil {
return fmt.Errorf("smtp data: %w", err)
}
_, err = w.Write(msg)
if err != nil {
return fmt.Errorf("write msg: %w", err)
}
err = w.Close()
if err != nil {
return fmt.Errorf("close writer: %w", err)
}
// Email is sent successfully after w.Close(), ignore Quit errors
// Some SMTP servers return non-standard responses on QUIT
_ = client.Quit()
return nil
}
// smtpTestRootCAs 仅供单元测试注入自签 CA,生产环境始终为 nil(走系统信任链)。
var smtpTestRootCAs *x509.CertPool
func smtpTLSConfig(host string) *tls.Config {
return &tls.Config{
ServerName: host,
// 强制 TLS 1.2+,避免协议降级导致的弱加密风险。
MinVersion: tls.VersionTLS12,
RootCAs: smtpTestRootCAs,
}
}
// connectSMTP 按配置建立 SMTP 会话,发送与测试连接共用此路径,
// 保证"测试连接成功 ⇔ 实际发信可用":
// - UseTLS=true:先尝试隐式 TLS(465 语义);若服务器以明文应答
// (587/25 等提交端口的 STARTTLS 语义),自动改走"明文连接 + 强制 STARTTLS"。
// 两种方式都无法建立加密连接时报错,绝不明文继续。
// - UseTLS=false:明文连接后若服务器支持 STARTTLS 则机会式升级,
// 与 smtp.SendMail 的默认行为一致。
func (s *EmailService) connectSMTP(config *SMTPConfig) (*smtp.Client, error) {
addr := fmt.Sprintf("%s:%d", config.Host, config.Port)
dialer := &net.Dialer{Timeout: smtpDialTimeout}
tlsConfig := smtpTLSConfig(config.Host)
if config.UseTLS {
conn, err := tls.DialWithDialer(dialer, "tcp", addr, tlsConfig)
if err == nil {
return newSMTPClient(conn, config.Host)
}
var recordErr tls.RecordHeaderError
if !errors.As(err, &recordErr) {
return nil, fmt.Errorf("tls dial: %w", err)
}
// SMTP 服务器先发问候语:明文问候会让 TLS 握手立刻返回
// RecordHeaderError,据此可靠判定对端期望 STARTTLS。
return s.connectSMTPStartTLS(dialer, addr, config.Host, tlsConfig, true)
}
return s.connectSMTPStartTLS(dialer, addr, config.Host, tlsConfig, false)
}
// connectSMTPStartTLS 建立明文连接并按需升级 STARTTLS。
// mandatory 为 true 时服务器必须支持 STARTTLS,否则报错。
func (s *EmailService) connectSMTPStartTLS(dialer *net.Dialer, addr, host string, tlsConfig *tls.Config, mandatory bool) (*smtp.Client, error) {
conn, err := dialer.Dial("tcp", addr)
if err != nil {
return nil, fmt.Errorf("smtp dial: %w", err)
}
client, err := newSMTPClient(conn, host)
if err != nil {
return nil, err
}
if ok, _ := client.Extension("STARTTLS"); !ok {
if mandatory {
_ = client.Close()
return nil, errors.New("smtp server does not support STARTTLS")
}
return client, nil
}
if err := client.StartTLS(tlsConfig); err != nil {
_ = client.Close()
return nil, fmt.Errorf("starttls: %w", err)
}
return client, nil
}
func newSMTPClient(conn net.Conn, host string) (*smtp.Client, error) {
_ = conn.SetDeadline(time.Now().Add(smtpIOTimeout))
client, err := smtp.NewClient(conn, host)
if err != nil {
_ = conn.Close()
return nil, fmt.Errorf("new smtp client: %w", err)
}
return client, nil
}
// GenerateVerifyCode 生成6位数字验证码
func (s *EmailService) GenerateVerifyCode() (string, error) {
const digits = "0123456789"
@@ -451,49 +449,24 @@ func (s *EmailService) buildVerifyCodeEmailBody(code, siteName string) string {
`, html.EscapeString(siteName), code)
}
// TestSMTPConnectionWithConfig 使用指定配置测试SMTP连接
// TestSMTPConnectionWithConfig 使用指定配置测试SMTP连接。
// 与 SendEmailWithConfig 共用 connectSMTP 建连(含 STARTTLS 升级逻辑),
// 避免出现"测试连接失败但实际发信成功"的不一致。
func (s *EmailService) TestSMTPConnectionWithConfig(config *SMTPConfig) error {
addr := fmt.Sprintf("%s:%d", config.Host, config.Port)
if config.UseTLS {
tlsConfig := &tls.Config{
ServerName: config.Host,
// 与发送逻辑一致,显式要求 TLS 1.2+。
MinVersion: tls.VersionTLS12,
}
conn, err := tls.Dial("tcp", addr, tlsConfig)
if err != nil {
return fmt.Errorf("tls connection failed: %w", err)
}
defer func() { _ = conn.Close() }()
client, err := smtp.NewClient(conn, config.Host)
if err != nil {
return fmt.Errorf("smtp client creation failed: %w", err)
}
defer func() { _ = client.Close() }()
auth := smtp.PlainAuth("", config.Username, config.Password, config.Host)
if err = client.Auth(auth); err != nil {
return fmt.Errorf("smtp authentication failed: %w", err)
}
return client.Quit()
}
// 非TLS连接测试
client, err := smtp.Dial(addr)
client, err := s.connectSMTP(config)
if err != nil {
return fmt.Errorf("smtp connection failed: %w", err)
}
defer func() { _ = client.Close() }()
auth := smtp.PlainAuth("", config.Username, config.Password, config.Host)
if err = client.Auth(auth); err != nil {
if err := client.Auth(auth); err != nil {
return fmt.Errorf("smtp authentication failed: %w", err)
}
return client.Quit()
// 认证成功即视为连接可用;与发送路径一致,忽略 QUIT 的非标准响应。
_ = client.Quit()
return nil
}
// GeneratePasswordResetToken generates a secure 32-byte random token (64 hex characters)
@@ -0,0 +1,383 @@
//go:build unit
package service
import (
"bufio"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"math/big"
"net"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
)
// newSMTPTestCert 生成 127.0.0.1/localhost 的自签证书及其信任池。
func newSMTPTestCert(t *testing.T) (tls.Certificate, *x509.CertPool) {
t.Helper()
priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("generate key: %v", err)
}
template := x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "127.0.0.1"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
BasicConstraintsValid: true,
IsCA: true,
IPAddresses: []net.IP{net.ParseIP("127.0.0.1")},
DNSNames: []string{"localhost"},
}
der, err := x509.CreateCertificate(rand.Reader, &template, &template, &priv.PublicKey, priv)
if err != nil {
t.Fatalf("create certificate: %v", err)
}
leaf, err := x509.ParseCertificate(der)
if err != nil {
t.Fatalf("parse certificate: %v", err)
}
pool := x509.NewCertPool()
pool.AddCert(leaf)
return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: priv}, pool
}
// fakeSMTPServer 是覆盖三种连接形态的最小 SMTP 服务器:
// 隐式 TLS(465 语义)、明文+STARTTLS(587 语义)、纯明文。
type fakeSMTPServer struct {
listener net.Listener
tlsConfig *tls.Config
advertiseStartTLS bool
mu sync.Mutex
commands []string
conns atomic.Int64
wg sync.WaitGroup
}
func startFakeSMTPServer(t *testing.T, implicitTLS, advertiseStartTLS bool) (*fakeSMTPServer, int) {
t.Helper()
cert, pool := newSMTPTestCert(t)
prevPool := smtpTestRootCAs
smtpTestRootCAs = pool
t.Cleanup(func() { smtpTestRootCAs = prevPool })
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen: %v", err)
}
srv := &fakeSMTPServer{
listener: listener,
tlsConfig: &tls.Config{Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12},
advertiseStartTLS: advertiseStartTLS,
}
if implicitTLS {
srv.listener = tls.NewListener(listener, srv.tlsConfig)
}
t.Cleanup(func() {
_ = srv.listener.Close()
srv.wg.Wait()
})
srv.wg.Add(1)
go func() {
defer srv.wg.Done()
for {
conn, err := srv.listener.Accept()
if err != nil {
return
}
srv.conns.Add(1)
srv.wg.Add(1)
go func() {
defer srv.wg.Done()
defer func() { _ = conn.Close() }()
_ = conn.SetDeadline(time.Now().Add(10 * time.Second))
srv.serve(conn, srv.advertiseStartTLS)
}()
}
}()
port := listener.Addr().(*net.TCPAddr).Port
return srv, port
}
func (srv *fakeSMTPServer) record(cmd string) {
srv.mu.Lock()
defer srv.mu.Unlock()
srv.commands = append(srv.commands, cmd)
}
func (srv *fakeSMTPServer) sawCommand(prefix string) bool {
srv.mu.Lock()
defer srv.mu.Unlock()
for _, cmd := range srv.commands {
if strings.HasPrefix(strings.ToUpper(cmd), prefix) {
return true
}
}
return false
}
func (srv *fakeSMTPServer) serve(conn net.Conn, allowStartTLS bool) {
reader := bufio.NewReader(conn)
writer := bufio.NewWriter(conn)
writeLine := func(line string) bool {
if _, err := writer.WriteString(line + "\r\n"); err != nil {
return false
}
return writer.Flush() == nil
}
if !writeLine("220 fake.test ESMTP ready") {
return
}
for {
line, err := reader.ReadString('\n')
if err != nil {
return
}
cmd := strings.TrimSpace(line)
srv.record(cmd)
upper := strings.ToUpper(cmd)
switch {
case strings.HasPrefix(upper, "EHLO"), strings.HasPrefix(upper, "HELO"):
ok := writeLine("250-fake.test")
if allowStartTLS {
ok = ok && writeLine("250-STARTTLS")
}
if !(ok && writeLine("250-AUTH PLAIN LOGIN") && writeLine("250 8BITMIME")) {
return
}
case upper == "STARTTLS" && allowStartTLS:
if !writeLine("220 2.0.0 ready to start TLS") {
return
}
tlsConn := tls.Server(conn, srv.tlsConfig)
if err := tlsConn.Handshake(); err != nil {
return
}
srv.serveUpgraded(tlsConn)
return
case strings.HasPrefix(upper, "AUTH"):
if !writeLine("235 2.7.0 authentication successful") {
return
}
case strings.HasPrefix(upper, "MAIL"), strings.HasPrefix(upper, "RCPT"):
if !writeLine("250 ok") {
return
}
case upper == "DATA":
if !writeLine("354 go ahead") {
return
}
for {
dataLine, err := reader.ReadString('\n')
if err != nil {
return
}
if strings.TrimRight(dataLine, "\r\n") == "." {
break
}
}
if !writeLine("250 message accepted") {
return
}
case upper == "QUIT":
_ = writeLine("221 bye")
return
default:
if !writeLine("250 ok") {
return
}
}
}
}
// serveUpgraded 复用命令循环处理 STARTTLS 升级后的会话(升级后不再提供 STARTTLS)。
func (srv *fakeSMTPServer) serveUpgraded(conn net.Conn) {
reader := bufio.NewReader(conn)
writer := bufio.NewWriter(conn)
// net/smtp 在 StartTLS 成功后会重新发送 EHLO,直接进入命令循环即可。
srv.serveCommands(reader, writer)
}
func (srv *fakeSMTPServer) serveCommands(reader *bufio.Reader, writer *bufio.Writer) {
writeLine := func(line string) bool {
if _, err := writer.WriteString(line + "\r\n"); err != nil {
return false
}
return writer.Flush() == nil
}
for {
line, err := reader.ReadString('\n')
if err != nil {
return
}
cmd := strings.TrimSpace(line)
srv.record(cmd)
upper := strings.ToUpper(cmd)
switch {
case strings.HasPrefix(upper, "EHLO"), strings.HasPrefix(upper, "HELO"):
if !(writeLine("250-fake.test") && writeLine("250-AUTH PLAIN LOGIN") && writeLine("250 8BITMIME")) {
return
}
case strings.HasPrefix(upper, "AUTH"):
if !writeLine("235 2.7.0 authentication successful") {
return
}
case strings.HasPrefix(upper, "MAIL"), strings.HasPrefix(upper, "RCPT"):
if !writeLine("250 ok") {
return
}
case upper == "DATA":
if !writeLine("354 go ahead") {
return
}
for {
dataLine, err := reader.ReadString('\n')
if err != nil {
return
}
if strings.TrimRight(dataLine, "\r\n") == "." {
break
}
}
if !writeLine("250 message accepted") {
return
}
case upper == "QUIT":
_ = writeLine("221 bye")
return
default:
if !writeLine("250 ok") {
return
}
}
}
}
func smtpTestConfig(port int, useTLS bool) *SMTPConfig {
return &SMTPConfig{
Host: "127.0.0.1",
Port: port,
Username: "user",
Password: "pass",
From: "noreply@example.com",
FromName: "Test",
UseTLS: useTLS,
}
}
// 465 语义:UseTLS=true + 隐式 TLS 服务器,原有路径保持可用。
func TestSMTPConnectionImplicitTLS(t *testing.T) {
srv, port := startFakeSMTPServer(t, true, false)
svc := &EmailService{}
if err := svc.TestSMTPConnectionWithConfig(smtpTestConfig(port, true)); err != nil {
t.Fatalf("expected implicit TLS connection to succeed, got: %v", err)
}
if !srv.sawCommand("EHLO") {
t.Fatal("expected server to receive EHLO")
}
}
// 587 语义(#1470/#1488 核心场景):UseTLS=true + 明文问候的 STARTTLS 服务器,
// 隐式 TLS 失败后必须自动降级为强制 STARTTLS 并成功。
func TestSMTPConnectionStartTLSFallbackWhenTLSEnabled(t *testing.T) {
srv, port := startFakeSMTPServer(t, false, true)
svc := &EmailService{}
if err := svc.TestSMTPConnectionWithConfig(smtpTestConfig(port, true)); err != nil {
t.Fatalf("expected STARTTLS fallback to succeed, got: %v", err)
}
if !srv.sawCommand("STARTTLS") {
t.Fatal("expected server to receive STARTTLS command")
}
if got := srv.conns.Load(); got < 2 {
t.Fatalf("expected implicit TLS attempt before STARTTLS fallback (>=2 connections), got %d", got)
}
}
// UseTLS=true 但服务器不支持 STARTTLS:必须报错,且绝不能把凭据发到明文连接上。
func TestSMTPConnectionMandatoryStartTLSRefusesPlaintext(t *testing.T) {
srv, port := startFakeSMTPServer(t, false, false)
svc := &EmailService{}
err := svc.TestSMTPConnectionWithConfig(smtpTestConfig(port, true))
if err == nil {
t.Fatal("expected error when server does not support STARTTLS")
}
if !strings.Contains(err.Error(), "STARTTLS") {
t.Fatalf("expected STARTTLS-related error, got: %v", err)
}
if srv.sawCommand("AUTH") {
t.Fatal("credentials must not be sent over plaintext when TLS is required")
}
}
// UseTLS=false + 服务器支持 STARTTLS:测试连接与发送路径一致,机会式升级后认证成功。
// 这是 #1488 评论"测试连接不成功,发送测试邮件实际上能发"的回归用例。
func TestSMTPConnectionOpportunisticStartTLSWhenTLSDisabled(t *testing.T) {
srv, port := startFakeSMTPServer(t, false, true)
svc := &EmailService{}
if err := svc.TestSMTPConnectionWithConfig(smtpTestConfig(port, false)); err != nil {
t.Fatalf("expected opportunistic STARTTLS test connection to succeed, got: %v", err)
}
if !srv.sawCommand("STARTTLS") {
t.Fatal("expected test connection to upgrade via STARTTLS like the send path")
}
}
// UseTLS=false + 服务器不支持 STARTTLS:保持明文直连(既有行为不回归)。
func TestSMTPConnectionPlainWhenNoStartTLS(t *testing.T) {
srv, port := startFakeSMTPServer(t, false, false)
svc := &EmailService{}
if err := svc.TestSMTPConnectionWithConfig(smtpTestConfig(port, false)); err != nil {
t.Fatalf("expected plain connection to succeed, got: %v", err)
}
if srv.sawCommand("STARTTLS") {
t.Fatal("did not expect STARTTLS command when server does not advertise it")
}
}
// 发送路径全流程:UseTLS=true + STARTTLS 服务器(587 语义)完整走完 MAIL/RCPT/DATA。
func TestSendEmailWithConfigStartTLSFallback(t *testing.T) {
srv, port := startFakeSMTPServer(t, false, true)
svc := &EmailService{}
err := svc.SendEmailWithConfig(smtpTestConfig(port, true), "rcpt@example.com", "subject", "<p>body</p>")
if err != nil {
t.Fatalf("expected send via STARTTLS fallback to succeed, got: %v", err)
}
if !srv.sawCommand("STARTTLS") {
t.Fatal("expected send path to upgrade via STARTTLS")
}
if !srv.sawCommand("DATA") {
t.Fatal("expected send path to reach DATA")
}
}
// 发送路径全流程:UseTLS=true + 隐式 TLS 服务器(465 语义)保持既有行为。
func TestSendEmailWithConfigImplicitTLS(t *testing.T) {
srv, port := startFakeSMTPServer(t, true, false)
svc := &EmailService{}
err := svc.SendEmailWithConfig(smtpTestConfig(port, true), "rcpt@example.com", "subject", "<p>body</p>")
if err != nil {
t.Fatalf("expected send via implicit TLS to succeed, got: %v", err)
}
if !srv.sawCommand("DATA") {
t.Fatal("expected send path to reach DATA")
}
}