From cf0cba0679d82e3575e4cf3b66958d86700991be Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 2 Sep 2026 21:59:47 +0800 Subject: [PATCH] refactor(mail): modernize smtp sending with go-mail and unify message gateway pusher --- .gitignore | 1 + backend/core/contracts/push.go | 3 + backend/core/extpoints/config_resolve.go | 2 +- backend/core/extpoints/config_value.go | 4 +- backend/go.mod | 6 +- backend/go.sum | 2 + backend/pkg/mail/errs.go | 11 +- backend/pkg/mail/mail.go | 322 +++++++----------- backend/pkg/mail/mail_test.go | 170 +++++---- .../domain/message_gateway/push/email.go | 58 +--- .../domain/message_gateway/push/email_test.go | 61 +++- 11 files changed, 298 insertions(+), 342 deletions(-) diff --git a/.gitignore b/.gitignore index 668bfcc9..7c44a6f4 100644 --- a/.gitignore +++ b/.gitignore @@ -67,3 +67,4 @@ s3_cache /backend/plugins/domain/upload/task/uploads/ /backend/data/ /backend/plugins/drivers/driver_http/dist/ +/backend/uploads/ diff --git a/backend/core/contracts/push.go b/backend/core/contracts/push.go index 9769a144..37863342 100644 --- a/backend/core/contracts/push.go +++ b/backend/core/contracts/push.go @@ -6,6 +6,7 @@ package contracts import "context" +// PushNotificationTemplate defines notification message template payload. type PushNotificationTemplate struct { Title string Content string @@ -13,6 +14,7 @@ type PushNotificationTemplate struct { Ext map[string]any } +// PushEventMeta defines metadata for a system push event. type PushEventMeta struct { Key string Name string @@ -20,6 +22,7 @@ type PushEventMeta struct { DefaultTemplate PushNotificationTemplate } +// PushRegistry defines the interface for registering built-in events. type PushRegistry interface { RegisterBuiltInEvent(meta PushEventMeta) SyncEvents(ctx context.Context) error diff --git a/backend/core/extpoints/config_resolve.go b/backend/core/extpoints/config_resolve.go index 3d8c7c37..4b2e1f7b 100644 --- a/backend/core/extpoints/config_resolve.go +++ b/backend/core/extpoints/config_resolve.go @@ -181,7 +181,7 @@ func formatEntryValue(value any, secret bool) string { return "" } rv := reflect.ValueOf(value) - if rv.Kind() == reflect.Ptr { + if rv.Kind() == reflect.Pointer { if rv.IsNil() { return "" } diff --git a/backend/core/extpoints/config_value.go b/backend/core/extpoints/config_value.go index b1cb0b98..5d1da6b5 100644 --- a/backend/core/extpoints/config_value.go +++ b/backend/core/extpoints/config_value.go @@ -33,7 +33,7 @@ func convertValue(raw any, typ reflect.Type) (any, error) { return convertSlice(raw, typ) case reflect.Struct: return convertStruct(raw, typ) - case reflect.Ptr: + case reflect.Pointer: return convertPointer(raw, typ) default: return nil, fmt.Errorf("%w: %s is not a supported configuration type", ErrConfigType, typ) @@ -44,7 +44,7 @@ func convertValue(raw any, typ reflect.Type) (any, error) { // Nested pointers are rejected so configuration tags stay one level deep. func convertPointer(raw any, typ reflect.Type) (any, error) { elemType := typ.Elem() - if elemType.Kind() == reflect.Ptr { + if elemType.Kind() == reflect.Pointer { return nil, fmt.Errorf("%w: %s is not a supported configuration type", ErrConfigType, typ) } elem, err := convertValue(raw, elemType) diff --git a/backend/go.mod b/backend/go.mod index 8a2303a5..5e85a362 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -16,8 +16,6 @@ require ( github.com/gin-contrib/sessions v1.1.0 github.com/gin-gonic/gin v1.12.0 github.com/glebarez/sqlite v1.11.0 - github.com/go-jose/go-jose/v4 v4.1.4 - github.com/google/go-cmp v0.7.0 github.com/google/uuid v1.6.0 github.com/gorilla/sessions v1.4.0 github.com/gorilla/websocket v1.5.3 @@ -38,6 +36,7 @@ require ( github.com/swaggo/swag v1.16.6 github.com/tencent-connect/botgo v0.2.1 github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2 + github.com/wneessen/go-mail v0.8.1 go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin v0.70.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.70.0 go.opentelemetry.io/otel v1.45.0 @@ -54,7 +53,6 @@ require ( gopkg.in/telebot.v4 v4.0.0-beta.10 gorm.io/driver/clickhouse v0.7.0 gorm.io/driver/postgres v1.6.2 - gorm.io/driver/sqlite v1.6.0 gorm.io/gorm v1.31.2 gorm.io/plugin/dbresolver v1.6.2 gorm.io/plugin/opentelemetry v0.1.14 @@ -95,6 +93,7 @@ require ( github.com/glebarez/go-sqlite v1.21.2 // indirect github.com/go-faster/city v1.0.1 // indirect github.com/go-faster/errors v0.7.1 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-logr/logr v1.4.4 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/go-openapi/jsonpointer v0.22.1 // indirect @@ -133,7 +132,6 @@ require ( github.com/klauspost/cpuid/v2 v2.4.0 // indirect github.com/leodido/go-urn v1.5.0 // indirect github.com/mattn/go-isatty v0.0.24 // indirect - github.com/mattn/go-sqlite3 v1.14.22 // indirect github.com/mfridman/interpolate v0.0.2 // indirect github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect github.com/modern-go/reflect2 v1.0.2 // indirect diff --git a/backend/go.sum b/backend/go.sum index 4b9b4830..5c8473c9 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -675,6 +675,8 @@ github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2 h1:3/aHKUq7qaFMWxyQV0W github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2/go.mod h1:Zit4b8AQXaXvA68+nzmbyDzqiyFRISyw1JiD5JqUBjw= github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2 h1:cj/Z6FKTTYBnstI0Lni9PA+k2foounKIPUmj1LBwNiQ= github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2/go.mod h1:LDaXk90gKEC2nC7JH3Lpnhfu+2V7o/TsqomJJmqA39o= +github.com/wneessen/go-mail v0.8.1 h1:tVcncj02/QySVFw3zr/kXOzZcuFQqBNT6K+Rbgm/pcM= +github.com/wneessen/go-mail v0.8.1/go.mod h1:dWZ61zadzCIyvB4y1/YzC5O7MrbbzBfPkARmbosdf8w= github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU= github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= github.com/yuin/goldmark v1.1.25/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= diff --git a/backend/pkg/mail/errs.go b/backend/pkg/mail/errs.go index 2c930984..caced8d3 100644 --- a/backend/pkg/mail/errs.go +++ b/backend/pkg/mail/errs.go @@ -5,12 +5,7 @@ package mail const ( - errDialTLSFailed = "dial tls failed: %w" - errSMTPClientCreationFailed = "smtp client creation failed: %w" - errSMTPAuthFailed = "smtp auth failed: %w" - errSMTPMailCommandFailed = "smtp mail command failed: %w" - errSMTPRcptCommandFailed = "smtp rcpt command failed: %w" - errSMTPDataCommandFailed = "smtp data command failed: %w" - errSMTPWritingBodyFailed = "smtp writing body failed: %w" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errSendMailFailed = "send mail failed: %w" + errCreateMailMessageFailed = "create mail message failed: %w" + errCreateMailClientFailed = "create mail client failed: %w" + errSendMailFailed = "send mail failed: %w" ) diff --git a/backend/pkg/mail/mail.go b/backend/pkg/mail/mail.go index b64ae5fe..2ae69df8 100644 --- a/backend/pkg/mail/mail.go +++ b/backend/pkg/mail/mail.go @@ -8,33 +8,39 @@ import ( "context" "crypto/tls" "fmt" - "net" - "net/smtp" - "strconv" "strings" "time" + + gomail "github.com/wneessen/go-mail" + golog "github.com/wneessen/go-mail/log" ) const ( - smtpSSLPort = 465 // SMTP SSL 端口 - smtpDialTimeout = 5 * time.Second // SMTP 连接超时 - smtpSessionDeadline = 10 * time.Second // SMTP 会话截止时间 + smtpSSLPort = 465 // SMTP SSL 端口 + smtpDialTimeout = 5 * time.Second // SMTP 连接超时 ) // Config represents SMTP mail configuration type Config struct { - Host string - Port int - Username string - Password string + Host string + Port int + Username string + Password string + FromName string // 可选发件人显示名称 + InsecureSkipVerify bool // 是否跳过证书校验 (默认跳过以兼容自签名证书) } -// sanitizeHeaderValue removes CR/LF bytes so untrusted values cannot inject -// additional email headers (email header injection). -func sanitizeHeaderValue(v string) string { - v = strings.ReplaceAll(v, "\r", "") - v = strings.ReplaceAll(v, "\n", "") - return v +// Option modifies internal mail send options +type Option func(*clientOptions) + +type clientOptions struct { + debugLogger golog.Logger +} + +func withLogger(l golog.Logger) Option { + return func(co *clientOptions) { + co.debugLogger = l + } } // SendMail sends an HTML email using the provided config and message details @@ -44,206 +50,112 @@ func SendMail(ctx context.Context, cfg Config, to, subject, body string) error { // SendMailHTML sends an HTML format email func SendMailHTML(ctx context.Context, cfg Config, to, subject, body string) error { - addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port)) + return send(ctx, cfg, to, subject, body) +} - // Header & MIME settings for HTML email - header := make(map[string]string) - header["From"] = sanitizeHeaderValue(cfg.Username) - header["To"] = sanitizeHeaderValue(to) - header["Subject"] = sanitizeHeaderValue(subject) - header["MIME-Version"] = "1.0" - header["Content-Type"] = "text/html; charset=UTF-8" +// SendMailWithLog sends a test email and records a detailed SMTP connection log +func SendMailWithLog(ctx context.Context, cfg Config, to, subject, body string) (string, error) { + var logBuf bytes.Buffer + logger := &bufferLogger{buf: &logBuf} - message := "" - for k, v := range header { - message += fmt.Sprintf("%s: %s\r\n", k, v) - } - message += "\r\n" + body + fmt.Fprintf(&logBuf, "[System] Connecting to %s:%d...\n", cfg.Host, cfg.Port) - auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host) - - // If using SSL port 465, we connection via TLS dial - if cfg.Port == smtpSSLPort { - return sendMailViaSSL(ctx, addr, auth, cfg, to, message) - } - - // For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it) - err := smtp.SendMail(addr, auth, cfg.Username, []string{to}, []byte(message)) + err := send(ctx, cfg, to, subject, body, withLogger(logger)) if err != nil { + fmt.Fprintf(&logBuf, "[Error] Mail sending failed: %v\n", err) + return logBuf.String(), err + } + + fmt.Fprintf(&logBuf, "[System] Mail sent successfully!\n") + return logBuf.String(), nil +} + +func send(ctx context.Context, cfg Config, to, subject, body string, opts ...Option) error { + var co clientOptions + for _, opt := range opts { + opt(&co) + } + + msg := gomail.NewMsg() + var err error + if cfg.FromName != "" { + err = msg.FromFormat(cfg.FromName, cfg.Username) + } else { + err = msg.From(cfg.Username) + } + if err != nil { + return fmt.Errorf(errCreateMailMessageFailed, err) + } + + if err = msg.To(to); err != nil { + return fmt.Errorf(errCreateMailMessageFailed, err) + } + + msg.Subject(subject) + msg.SetBodyString(gomail.TypeTextHTML, body) + + clientOpts := []gomail.Option{ + gomail.WithPort(cfg.Port), + gomail.WithTimeout(smtpDialTimeout), + } + + tlsConfig := &tls.Config{ + InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates + ServerName: cfg.Host, + } + clientOpts = append(clientOpts, gomail.WithTLSConfig(tlsConfig)) + + if cfg.Port == smtpSSLPort { + clientOpts = append(clientOpts, gomail.WithSSL()) + } else { + clientOpts = append(clientOpts, gomail.WithTLSPolicy(gomail.TLSOpportunistic)) + } + + if cfg.Username != "" && cfg.Password != "" { + clientOpts = append(clientOpts, + gomail.WithSMTPAuth(gomail.SMTPAuthPlain), + gomail.WithUsername(cfg.Username), + gomail.WithPassword(cfg.Password), + ) + } + + if co.debugLogger != nil { + clientOpts = append(clientOpts, gomail.WithDebugLog(), gomail.WithLogger(co.debugLogger)) + } + + client, err := gomail.NewClient(cfg.Host, clientOpts...) + if err != nil { + return fmt.Errorf(errCreateMailClientFailed, err) + } + defer func() { _ = client.Close() }() + + if err = client.DialAndSendWithContext(ctx, msg); err != nil { return fmt.Errorf(errSendMailFailed, err) } return nil } -// sendMailViaSSL 通过 TLS 直接连接 SMTP SSL 端口发送邮件 -func sendMailViaSSL(ctx context.Context, addr string, auth smtp.Auth, cfg Config, to, message string) error { - tlsConfig := &tls.Config{ - InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates - ServerName: cfg.Host, - } - dialer := &net.Dialer{Timeout: smtpDialTimeout} - tlsDialer := &tls.Dialer{ - NetDialer: dialer, - Config: tlsConfig, - } - conn, err := tlsDialer.DialContext(ctx, "tcp", addr) - if err != nil { - return fmt.Errorf(errDialTLSFailed, err) - } - defer func() { _ = conn.Close() }() - _ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline)) - - client, err := smtp.NewClient(conn, cfg.Host) - if err != nil { - return fmt.Errorf(errSMTPClientCreationFailed, err) - } - defer func() { _ = client.Close() }() - - if err = client.Auth(auth); err != nil { - return fmt.Errorf(errSMTPAuthFailed, err) - } - if err = client.Mail(cfg.Username); err != nil { - return fmt.Errorf(errSMTPMailCommandFailed, err) - } - if err = client.Rcpt(to); err != nil { - return fmt.Errorf(errSMTPRcptCommandFailed, err) - } - - w, err := client.Data() - if err != nil { - return fmt.Errorf(errSMTPDataCommandFailed, err) - } - defer func() { _ = w.Close() }() - - _, err = w.Write([]byte(message)) - if err != nil { - return fmt.Errorf(errSMTPWritingBodyFailed, err) - } - return nil +type bufferLogger struct { + buf *bytes.Buffer } -// SendMailWithLog sends a test email and records a detailed SMTP connection log -func SendMailWithLog(ctx context.Context, cfg Config, to, subject, body string) (string, error) { - var logBuf bytes.Buffer - logLine := func(dir, format string, args ...interface{}) { - fmt.Fprintf(&logBuf, "[%s] %s\n", dir, fmt.Sprintf(format, args...)) +func (l *bufferLogger) log(level string, entry golog.Log) { + msg := fmt.Sprintf(entry.Format, entry.Messages...) + msg = strings.TrimRight(msg, "\r\n") + var dir string + switch entry.Direction { + case golog.DirClientToServer: + dir = "C" + case golog.DirServerToClient: + dir = "S" + default: + dir = level } - - addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port)) - logLine("System", "Connecting to %s...", addr) - - var conn net.Conn - var err error - dialer := &net.Dialer{Timeout: smtpDialTimeout} - if cfg.Port == smtpSSLPort { - tlsConfig := &tls.Config{ - InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates - ServerName: cfg.Host, - } - tlsDialer := &tls.Dialer{ - NetDialer: dialer, - Config: tlsConfig, - } - conn, err = tlsDialer.DialContext(ctx, "tcp", addr) - } else { - conn, err = dialer.DialContext(ctx, "tcp", addr) - } - if err != nil { - logLine("Error", "Connection failed: %v", err) - return logBuf.String(), err - } - defer func() { _ = conn.Close() }() - logLine("System", "Connected successfully.") - - // Set a 10-second session deadline for read/write operations - _ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline)) - - client, err := smtp.NewClient(conn, cfg.Host) - if err != nil { - logLine("Error", "SMTP client handshake failed: %v", err) - return logBuf.String(), err - } - defer func() { _ = client.Close() }() - - // If not 465, support STARTTLS if available - if cfg.Port != smtpSSLPort { - if ok, _ := client.Extension("STARTTLS"); ok { - logLine("C", "STARTTLS") - tlsConfig := &tls.Config{ - InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates - ServerName: cfg.Host, - } - if err = client.StartTLS(tlsConfig); err != nil { - logLine("Error", "STARTTLS failed: %v", err) - return logBuf.String(), err - } - logLine("S", "220 Ready to start TLS") - } - } - - // Authentication - if cfg.Username != "" && cfg.Password != "" { - auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host) - logLine("C", "AUTH PLAIN **********") - if err = client.Auth(auth); err != nil { - logLine("Error", "Authentication failed: %v", err) - return logBuf.String(), err - } - logLine("S", "235 Authentication successful") - } - - // Mail command - logLine("C", "MAIL FROM:<%s>", cfg.Username) - if err = client.Mail(cfg.Username); err != nil { - logLine("Error", "MAIL FROM command failed: %v", err) - return logBuf.String(), err - } - logLine("S", "250 OK") - - // Rcpt command - logLine("C", "RCPT TO:<%s>", to) - if err = client.Rcpt(to); err != nil { - logLine("Error", "RCPT TO command failed: %v", err) - return logBuf.String(), err - } - logLine("S", "250 OK") - - // Data command - logLine("C", "DATA") - w, err := client.Data() - if err != nil { - logLine("Error", "DATA command failed: %v", err) - return logBuf.String(), err - } - logLine("S", "354 Start mail input") - - // Header & MIME settings for HTML email - header := make(map[string]string) - header["From"] = sanitizeHeaderValue(cfg.Username) - header["To"] = sanitizeHeaderValue(to) - header["Subject"] = sanitizeHeaderValue(subject) - header["MIME-Version"] = "1.0" - header["Content-Type"] = "text/html; charset=UTF-8" - - message := "" - for k, v := range header { - message += fmt.Sprintf("%s: %s\r\n", k, v) - } - message += "\r\n" + body - - logLine("System", "Sending message body...") - if _, err = w.Write([]byte(message)); err != nil { - _ = w.Close() - logLine("Error", "Writing message body failed: %v", err) - return logBuf.String(), err - } - _ = w.Close() - logLine("S", "250 OK") - - logLine("C", "QUIT") - _ = client.Quit() - logLine("System", "Mail sent successfully!") - - return logBuf.String(), nil + fmt.Fprintf(l.buf, "[%s] %s\n", dir, msg) } + +func (l *bufferLogger) Debugf(e golog.Log) { l.log("Debug", e) } +func (l *bufferLogger) Infof(e golog.Log) { l.log("Info", e) } +func (l *bufferLogger) Warnf(e golog.Log) { l.log("Warn", e) } +func (l *bufferLogger) Errorf(e golog.Log) { l.log("Error", e) } diff --git a/backend/pkg/mail/mail_test.go b/backend/pkg/mail/mail_test.go index 08e16330..4635292a 100644 --- a/backend/pkg/mail/mail_test.go +++ b/backend/pkg/mail/mail_test.go @@ -8,74 +8,105 @@ import ( "context" "net" "net/textproto" + "strings" "testing" ) -func TestSendMailMock(t *testing.T) { - // Start a mock SMTP server +func startMockSMTPServer(t *testing.T) (int, func()) { + t.Helper() l, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("failed to start mock smtp server: %v", err) } - defer func() { _ = l.Close() }() port := l.Addr().(*net.TCPAddr).Port go func() { - conn, err := l.Accept() + for { + conn, err := l.Accept() + if err != nil { + return + } + go handleMockSMTPConn(conn) + } + }() + + return port, func() { _ = l.Close() } +} + +func handleMockSMTPConn(conn net.Conn) { + defer func() { _ = conn.Close() }() + + writer := bufio.NewWriter(conn) + reader := bufio.NewReader(conn) + tp := textproto.NewReader(reader) + + // 220 Ready + _, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n") + _ = writer.Flush() + + for { + line, err := tp.ReadLine() if err != nil { return } - defer func() { _ = conn.Close() }() - - writer := bufio.NewWriter(conn) - reader := bufio.NewReader(conn) - tp := textproto.NewReader(reader) - - // 220 Ready - _, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n") - _ = writer.Flush() - - // Read HELO/EHLO - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n") - _ = writer.Flush() - - // Read AUTH PLAIN - _, _ = tp.ReadLine() - _, _ = writer.WriteString("235 Authentication successful\r\n") - _ = writer.Flush() - - // Read MAIL FROM - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() - - // Read RCPT TO - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() - - // Read DATA - _, _ = tp.ReadLine() - _, _ = writer.WriteString("354 Start mail input\r\n") - _ = writer.Flush() - - // Read body lines until dot - for { - line, err := tp.ReadLine() - if err != nil || line == "." { - break + upper := strings.ToUpper(line) + switch { + case strings.HasPrefix(upper, "EHLO") || strings.HasPrefix(upper, "HELO"): + _, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "AUTH PLAIN"): + _, _ = writer.WriteString("235 2.7.0 Authentication successful\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "MAIL FROM:"): + _, _ = writer.WriteString("250 2.1.0 Ok\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "RCPT TO:"): + _, _ = writer.WriteString("250 2.1.5 Ok\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "DATA"): + _, _ = writer.WriteString("354 Start mail input; end with .\r\n") + _ = writer.Flush() + for { + dataLine, err := tp.ReadLine() + if err != nil || dataLine == "." { + break + } } + _, _ = writer.WriteString("250 2.0.0 Ok: queued\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "QUIT"): + _, _ = writer.WriteString("221 2.0.0 Bye\r\n") + _ = writer.Flush() + return + default: + _, _ = writer.WriteString("250 Ok\r\n") + _ = writer.Flush() } - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() + } +} - // Read QUIT - _, _ = tp.ReadLine() - _, _ = writer.WriteString("221 Bye\r\n") - _ = writer.Flush() - }() +func TestSendMailMock(t *testing.T) { + port, cleanup := startMockSMTPServer(t) + defer cleanup() + + cfg := Config{ + Host: "127.0.0.1", + Port: port, + Username: "test@example.com", + Password: "password", + FromName: "Wavelet Notifier", + } + + err := SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "

Test Body

") + if err != nil { + t.Fatalf("failed to send mail: %v", err) + } +} + +func TestSendMailWithLog(t *testing.T) { + port, cleanup := startMockSMTPServer(t) + defer cleanup() cfg := Config{ Host: "127.0.0.1", @@ -84,28 +115,29 @@ func TestSendMailMock(t *testing.T) { Password: "password", } - err = SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "

Test Body

") + logs, err := SendMailWithLog(context.Background(), cfg, "recipient@example.com", "Test Subject", "

Test Log

") if err != nil { - t.Errorf("failed to send mail: %v", err) + t.Fatalf("failed to send mail with log: %v, log output:\n%s", err, logs) + } + + if !strings.Contains(logs, "[System] Connecting to") { + t.Errorf("expected connection log in output, got: %s", logs) + } + if !strings.Contains(logs, "[System] Mail sent successfully!") { + t.Errorf("expected success log in output, got: %s", logs) } } -func TestSanitizeHeaderValue(t *testing.T) { - tests := []struct { - name string - input string - want string - }{ - {"plain", "System Notification", "System Notification"}, - {"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"}, - {"cr stripped", "a\rb", "ab"}, - {"lf stripped", "a\nb", "ab"}, +func TestSendMailInvalidAddress(t *testing.T) { + cfg := Config{ + Host: "127.0.0.1", + Port: 25, + Username: "test@example.com", + Password: "password", } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := sanitizeHeaderValue(tt.input); got != tt.want { - t.Errorf("sanitizeHeaderValue(%q) = %q, want %q", tt.input, got, tt.want) - } - }) + + err := SendMail(context.Background(), cfg, "invalid address with \n newline", "Subject", "Body") + if err == nil { + t.Errorf("expected error for invalid address, got nil") } } diff --git a/backend/plugins/domain/message_gateway/push/email.go b/backend/plugins/domain/message_gateway/push/email.go index 2948c577..6c672ae6 100644 --- a/backend/plugins/domain/message_gateway/push/email.go +++ b/backend/plugins/domain/message_gateway/push/email.go @@ -4,30 +4,21 @@ package push import ( - "Wavelet/pkg/util" + pkgmail "Wavelet/pkg/mail" "context" "errors" "fmt" "net" - "net/smtp" - "strings" + "strconv" ) func init() { Register("email", &EmailPusher{}) } -// EmailPusher 极简 SMTP 邮件推送实现 (静态、解耦) +// EmailPusher 基于 pkg/mail 的 SMTP 邮件推送实现 type EmailPusher struct{} -// sanitizeEmailHeader removes CR/LF bytes so untrusted values cannot inject -// additional email headers (email header injection). -func sanitizeEmailHeader(v string) string { - v = strings.ReplaceAll(v, "\r", "") - v = strings.ReplaceAll(v, "\n", "") - return v -} - // Send 发送邮件 func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) (string, error) { if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" { @@ -40,11 +31,6 @@ func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body title := bodyTitle(body) content := bodyContent(body, "

%s: %v

", "") - // 邮件头和体 - from := cfg.Key - to := target - - // 如果 ext 中指定了 from_name,我们在 From 头部包含它 fromName := "System Notification" if ext != nil { if fn, ok := ext["from_name"].(string); ok && fn != "" { @@ -52,38 +38,26 @@ func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body } } - subjectHeader := fmt.Sprintf("Subject: %s\r\n", sanitizeEmailHeader(title)) - fromHeader := fmt.Sprintf("From: %s <%s>\r\n", sanitizeEmailHeader(fromName), sanitizeEmailHeader(from)) - toHeader := fmt.Sprintf("To: %s\r\n", sanitizeEmailHeader(to)) - mimeHeader := "MIME-version: 1.0;\r\nContent-Type: text/html; charset=\"UTF-8\";\r\n\r\n" - - // 拼装完整的邮件报文 - // 简单的 HTML 正文渲染 htmlBody := fmt.Sprintf(`

%s

%s
`, title, content) - msg := []byte(fromHeader + toHeader + subjectHeader + mimeHeader + htmlBody + "\r\n") - // 解析 Host 和 Port - host, port, err := net.SplitHostPort(cfg.URL) + host, portStr, err := net.SplitHostPort(cfg.URL) + port := 25 if err != nil { host = cfg.URL - port = "25" // 默认 SMTP 端口 + } else if p, err := strconv.Atoi(portStr); err == nil && p > 0 { + port = p } - auth := smtp.PlainAuth("", cfg.Key, cfg.Secret, host) + mailCfg := pkgmail.Config{ + Host: host, + Port: port, + Username: cfg.Key, + Password: cfg.Secret, + FromName: fromName, + } - // 异步超时处理 - errChan := make(chan error, 1) - util.Go(func() { - errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg) - }) - - select { - case <-ctx.Done(): - return "", ctx.Err() - case err := <-errChan: - if err != nil { - return "", fmt.Errorf("email: send smtp mail failed: %w", err) - } + if err := pkgmail.SendMail(ctx, mailCfg, target, title, htmlBody); err != nil { + return "", fmt.Errorf("email: send smtp mail failed: %w", err) } return "", nil diff --git a/backend/plugins/domain/message_gateway/push/email_test.go b/backend/plugins/domain/message_gateway/push/email_test.go index fb7367cc..3d97aac4 100644 --- a/backend/plugins/domain/message_gateway/push/email_test.go +++ b/backend/plugins/domain/message_gateway/push/email_test.go @@ -3,24 +3,63 @@ package push -import "testing" +import ( + "context" + "testing" +) + +func TestEmailPusherValidateConfig(t *testing.T) { + pusher := &EmailPusher{} -func TestSanitizeEmailHeader(t *testing.T) { tests := []struct { - name string - input string - want string + name string + cfg Config + wantErr bool }{ - {"plain", "System Notification", "System Notification"}, - {"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"}, - {"cr stripped", "a\rb", "ab"}, - {"lf stripped", "a\nb", "ab"}, + { + name: "empty url", + cfg: Config{URL: "", Key: "user", Secret: "pass"}, + wantErr: true, + }, + { + name: "empty key", + cfg: Config{URL: "smtp.example.com:587", Key: "", Secret: "pass"}, + wantErr: true, + }, + { + name: "empty secret", + cfg: Config{URL: "smtp.example.com:587", Key: "user", Secret: ""}, + wantErr: true, + }, + { + name: "valid config", + cfg: Config{URL: "smtp.example.com:587", Key: "user", Secret: "pass"}, + wantErr: false, + }, } + for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := sanitizeEmailHeader(tt.input); got != tt.want { - t.Errorf("sanitizeEmailHeader(%q) = %q, want %q", tt.input, got, tt.want) + err := pusher.ValidateConfig(tt.cfg) + if (err != nil) != tt.wantErr { + t.Errorf("ValidateConfig() error = %v, wantErr %v", err, tt.wantErr) } }) } } + +func TestEmailPusherSendValidation(t *testing.T) { + pusher := &EmailPusher{} + + // Missing target + _, err := pusher.Send(context.Background(), Config{URL: "127.0.0.1:25", Key: "u", Secret: "p"}, "", map[string]any{"title": "hi"}, "", nil) + if err == nil { + t.Errorf("expected error for empty target, got nil") + } + + // Missing config + _, err = pusher.Send(context.Background(), Config{}, "test@example.com", map[string]any{"title": "hi"}, "", nil) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } +}