feat: enhance email validation to prevent injection attacks and improve related tests

This commit is contained in:
MengMengCode
2026-08-11 19:41:06 +08:00
parent ede7a8aa19
commit 30afe090d3
3 changed files with 67 additions and 17 deletions
+33 -2
View File
@@ -28,11 +28,23 @@ func writePlainTextMail(
if strings.ContainsAny(subject, "\r\n\x00") {
return errors.New("email subject contains a prohibited control character")
}
fromHeader, err := validatedMailHeaderAddress(from)
if err != nil {
return fmt.Errorf("invalid email sender: %w", err)
}
recipientHeaders := make([]string, 0, len(recipients))
for _, recipient := range recipients {
header, err := validatedMailHeaderAddress(recipient)
if err != nil {
return fmt.Errorf("invalid email recipient: %w", err)
}
recipientHeaders = append(recipientHeaders, header)
}
encodedBody := wrapMIMEBase64(base64.StdEncoding.EncodeToString([]byte(body)))
message := strings.Join([]string{
"Date: " + time.Now().UTC().Format(time.RFC1123Z),
"From: " + formatMailAddress(from),
"To: " + joinMailAddresses(recipients),
"From: " + fromHeader,
"To: " + strings.Join(recipientHeaders, ", "),
"Subject: " + mime.QEncoding.Encode("UTF-8", subject),
"MIME-Version: 1.0",
"Content-Type: text/plain; charset=UTF-8",
@@ -52,6 +64,25 @@ func writePlainTextMail(
return nil
}
// validatedMailHeaderAddress keeps writePlainTextMail safe even if a future
// caller constructs mail.Address directly instead of using parseMailAddress.
func validatedMailHeaderAddress(address *mail.Address) (string, error) {
if address == nil || address.Address == "" || strings.TrimSpace(address.Address) != address.Address ||
strings.ContainsAny(address.Address, "\r\n\x00") {
return "", errors.New("email address contains a prohibited control character")
}
parsed, err := mail.ParseAddress(address.Address)
if err != nil || parsed.Name != "" || parsed.Address != address.Address {
return "", errors.New("invalid email address")
}
for _, character := range address.Name {
if character < 0x20 || character == 0x7f {
return "", errors.New("email display name contains a prohibited control character")
}
}
return formatMailAddress(address), nil
}
func wrapMIMEBase64(value string) string {
if value == "" {
return ""
+31
View File
@@ -42,3 +42,34 @@ func TestWritePlainTextMailRejectsInjectedSubject(t *testing.T) {
t.Fatal("injected subject was accepted")
}
}
func TestWritePlainTextMailRejectsDirectlyConstructedInjectedAddresses(t *testing.T) {
tests := []struct {
name string
from *mail.Address
recipients []*mail.Address
}{
{
name: "sender address",
from: &mail.Address{Address: "[email protected]\r\nBcc: [email protected]"},
recipients: []*mail.Address{{Address: "[email protected]"}},
},
{
name: "sender display name",
from: &mail.Address{Name: "Alerts\r\nBcc: [email protected]", Address: "[email protected]"},
recipients: []*mail.Address{{Address: "[email protected]"}},
},
{
name: "recipient address",
from: &mail.Address{Address: "[email protected]"},
recipients: []*mail.Address{{Address: "[email protected]\nCc: [email protected]"}},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if err := writePlainTextMail(&bytes.Buffer{}, test.from, test.recipients, "subject", "body"); err == nil {
t.Fatal("injected address was accepted")
}
})
}
}
+3 -15
View File
@@ -799,14 +799,10 @@ func sendEmailNotificationTest(ctx context.Context, config map[string]any) error
// Addresses are parsed as RFC mailboxes, the subject rejects control
// characters, and the body is MIME-base64 encoded by writePlainTextMail.
// CodeQL's email-injection query has no sanitizer model for these steps.
// Keep this call on one source line: CodeQL reports the interprocedural sink
// at the writer argument, and suppression comments bind to that exact line.
// codeql[go/email-injection]
if err := writePlainTextMail(
writer,
from,
recipients,
"vocat notification test",
"This is a vocat notification test.",
); err != nil {
if err := writePlainTextMail(writer, from, recipients, "vocat notification test", "This is a vocat notification test."); err != nil {
_ = writer.Close()
return fmt.Errorf("write SMTP test message: %w", err)
}
@@ -819,14 +815,6 @@ func sendEmailNotificationTest(ctx context.Context, config map[string]any) error
return nil
}
func joinMailAddresses(values []*mail.Address) string {
result := make([]string, 0, len(values))
for _, value := range values {
result = append(result, formatMailAddress(value))
}
return strings.Join(result, ", ")
}
func parseMailAddress(value string) (*mail.Address, error) {
value = strings.TrimSpace(value)
if value == "" || strings.ContainsAny(value, "\r\n\x00") {