Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 16 additions & 3 deletions internal/api/handler_registrations.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ func (h *Handler) submitRegistration(ctx context.Context, req *events.LambdaFunc
return nil, NewClientError(400, "invalid request body")
}

if err := validateRegistrationRequest(body); err != nil {
if err := validateRegistrationRequest(&body); err != nil {
return nil, err
}

Expand Down Expand Up @@ -394,8 +394,21 @@ func (h *Handler) deleteRegistration(ctx context.Context, id string) (any, error
return map[string]string{"status": "deleted"}, nil
}

// validateRegistrationRequest checks required fields for a registration submission.
func validateRegistrationRequest(req RegistrationRequest) error {
// maxAccountNameLen caps the persisted/displayed account name length as a
// defense-in-depth measure against abuse of the unauthenticated registration
// endpoint (#544 / #401).
const maxAccountNameLen = 256

// validateRegistrationRequest checks required fields for a registration
// submission. It also normalizes account_name in place: CR/LF characters are
// stripped (defense-in-depth against email header injection at the data
// source, complementing sanitizeHeader on the email subject path) and the
// result is length-capped.
func validateRegistrationRequest(req *RegistrationRequest) error {
req.AccountName = strings.TrimSpace(strings.NewReplacer("\r", "", "\n", "").Replace(req.AccountName))
if r := []rune(req.AccountName); len(r) > maxAccountNameLen {
req.AccountName = string(r[:maxAccountNameLen])
}
if !validAccountProviders[req.Provider] {
return NewClientError(400, "provider must be one of: aws, azure, gcp")
}
Expand Down
66 changes: 66 additions & 0 deletions internal/api/handler_registrations_validate_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
package api

import (
"strings"
"testing"
"unicode/utf8"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.

// TestValidateRegistrationRequest_AccountNameSanitized is the defense-in-depth
// regression test for #544 / #401: account_name from the unauthenticated
// POST /api/register endpoint must be CR/LF-stripped and length-capped at the
// data source before it is persisted or interpolated into any email header.
func TestValidateRegistrationRequest_AccountNameSanitized(t *testing.T) {
t.Run("strips CRLF", func(t *testing.T) {
req := RegistrationRequest{
Provider: "aws",
ExternalID: "ext-123",
AccountName: "Acme\r\nBcc: attacker@evil.example.com",
ContactEmail: "user@example.com",
}
require.NoError(t, validateRegistrationRequest(&req))
assert.NotContains(t, req.AccountName, "\r")
assert.NotContains(t, req.AccountName, "\n")
})

t.Run("length-capped (rune-safe)", func(t *testing.T) {
// Use a multibyte rune so the cap is exercised by rune count, not byte
// length, and confirm truncation never splits a rune (valid UTF-8).
req := RegistrationRequest{
Provider: "aws",
ExternalID: "ext-123",
AccountName: strings.Repeat("é", maxAccountNameLen+50),
ContactEmail: "user@example.com",
}
require.NoError(t, validateRegistrationRequest(&req))
assert.LessOrEqual(t, utf8.RuneCountInString(req.AccountName), maxAccountNameLen)
assert.True(t, utf8.ValidString(req.AccountName), "cap must not split a multibyte rune")
})

t.Run("CRLF-only name is rejected as empty", func(t *testing.T) {
req := RegistrationRequest{
Provider: "aws",
ExternalID: "ext-123",
AccountName: "\r\n",
ContactEmail: "user@example.com",
}
err := validateRegistrationRequest(&req)
require.Error(t, err)
assert.Contains(t, err.Error(), "account_name is required")
})

t.Run("whitespace-only name is rejected as empty", func(t *testing.T) {
req := RegistrationRequest{
Provider: "aws",
ExternalID: "ext-123",
AccountName: " \t ",
ContactEmail: "user@example.com",
}
err := validateRegistrationRequest(&req)
require.Error(t, err)
assert.Contains(t, err.Error(), "account_name is required")
})
}
7 changes: 6 additions & 1 deletion internal/email/templates.go
Original file line number Diff line number Diff line change
Expand Up @@ -805,7 +805,12 @@ func (s *Sender) SendRegistrationReceivedNotification(ctx context.Context, data
if err != nil {
return fmt.Errorf("failed to render registration received email: %w", err)
}
subject := fmt.Sprintf("CUDly - New Account Registration: %s (%s)", data.AccountName, data.Provider)
// sanitizeHeader strips CR/LF from the attacker-controlled AccountName /
// Provider (sourced from the unauthenticated POST /api/register endpoint)
// before they are interpolated into the SES email subject, preventing
// email header injection (#544 / #401). Mirrors the SMTP path fix.
subject := fmt.Sprintf("CUDly - New Account Registration: %s (%s)",
sanitizeHeader(data.AccountName), sanitizeHeader(data.Provider))
if data.RecipientEmail == "" {
return s.SendNotification(ctx, subject, body)
}
Expand Down
46 changes: 46 additions & 0 deletions internal/email/templates_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -613,3 +613,49 @@ func containsAnyStr(s string, subs ...string) bool {
}
return false
}

// TestSender_SendRegistrationReceivedNotification_SubjectHeaderInjection is the
// SES-path regression test for #544 / #401: a CR+LF in the attacker-controlled
// AccountName / Provider (sourced from the unauthenticated POST /api/register
// endpoint) must be stripped before the subject reaches the SES SendEmail API,
// so it cannot inject additional email headers. Mirrors the SMTP-path test.
func TestSender_SendRegistrationReceivedNotification_SubjectHeaderInjection(t *testing.T) {
mockSES := new(MockSESClient)
// Production mode so the send proceeds straight to SendEmail.
mockSES.On("GetAccount", mock.Anything, mock.AnythingOfType("*sesv2.GetAccountInput")).
Return(&sesv2.GetAccountOutput{ProductionAccessEnabled: true}, nil)

var capturedSubject string
mockSES.On("SendEmail", mock.Anything, mock.AnythingOfType("*sesv2.SendEmailInput")).
Run(func(args mock.Arguments) {
input := args.Get(1).(*sesv2.SendEmailInput)
capturedSubject = aws.ToString(input.Content.Simple.Subject.Data)
}).
Return(&sesv2.SendEmailOutput{MessageId: aws.String("msg-injection")}, nil)

sender := NewSenderWithClients(nil, mockSES, SenderConfig{
FromEmail: "noreply@example.com",
})

injectedName := "Acme\r\nBcc: attacker@evil.example.com"
injectedProvider := "aws\r\nX-Injected: yes"
data := RegistrationNotificationData{
AccountName: injectedName,
Provider: injectedProvider,
ExternalID: "ext-123",
ContactEmail: "registrant@example.com",
RecipientEmail: "admin@example.com",
}

ctx := context.Background()
err := sender.SendRegistrationReceivedNotification(ctx, data)
require.NoError(t, err)
mockSES.AssertExpectations(t)

// The subject that reached SES must contain no CR/LF characters.
assert.NotContains(t, capturedSubject, "\r", "SES subject must not contain CR: %q", capturedSubject)
assert.NotContains(t, capturedSubject, "\n", "SES subject must not contain LF: %q", capturedSubject)
// The injected header names must not survive into the subject as injectable headers.
assert.NotContains(t, capturedSubject, "\nBcc:")
assert.NotContains(t, capturedSubject, "\nX-Injected:")
}
Loading