From 4c80d160dd8d4bd56ff76cf64200925485187f94 Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 31 Jul 2026 23:12:02 +0800 Subject: [PATCH] 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 --- backend/internal/service/email_service.go | 211 +++++----- .../service/email_service_smtp_test.go | 383 ++++++++++++++++++ 2 files changed, 475 insertions(+), 119 deletions(-) create mode 100644 backend/internal/service/email_service_smtp_test.go diff --git a/backend/internal/service/email_service.go b/backend/internal/service/email_service.go index 0825257546..8e60d9e73d 100644 --- a/backend/internal/service/email_service.go +++ b/backend/internal/service/email_service.go @@ -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) diff --git a/backend/internal/service/email_service_smtp_test.go b/backend/internal/service/email_service_smtp_test.go new file mode 100644 index 0000000000..b9677a1af3 --- /dev/null +++ b/backend/internal/service/email_service_smtp_test.go @@ -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", "

body

") + 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", "

body

") + 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") + } +}