diff --git a/backend/internal/service/app_images_service.go b/backend/internal/service/app_images_service.go index f93918f9..e21d37b3 100644 --- a/backend/internal/service/app_images_service.go +++ b/backend/internal/service/app_images_service.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "fmt" "io" "mime/multipart" @@ -104,6 +105,9 @@ func (s *AppImagesService) UpdateImage(ctx context.Context, file *multipart.File defer fileReader.Close() strippedReader, err := imageutil.StripMetadata(fileReader, fileType) + if errors.Is(err, imageutil.ErrInvalidImage) { + return apperror.InvalidImage(err) + } if err != nil { return err } diff --git a/backend/internal/service/oidc_service.go b/backend/internal/service/oidc_service.go index d339c446..049610c5 100644 --- a/backend/internal/service/oidc_service.go +++ b/backend/internal/service/oidc_service.go @@ -515,6 +515,9 @@ func (s *OidcService) UpdateClientLogo(ctx context.Context, clientID string, fil } defer reader.Close() strippedReader, err := imageutil.StripMetadata(reader, fileType) + if errors.Is(err, imageutil.ErrInvalidImage) { + return apperror.InvalidImage(err) + } if err != nil { return err } diff --git a/backend/internal/utils/image/metadata.go b/backend/internal/utils/image/metadata.go index 9bf02dd0..77217df7 100644 --- a/backend/internal/utils/image/metadata.go +++ b/backend/internal/utils/image/metadata.go @@ -3,6 +3,8 @@ package profilepicture import ( "bytes" "encoding/binary" + "errors" + "fmt" "io" "strings" @@ -21,13 +23,10 @@ func StripMetadata(file io.Reader, ext string) (*bytes.Reader, error) { } switch strings.ToLower(ext) { - case "jpg", "jpeg": - stripped, err := exifremove.Remove(data) - if err == nil { - return bytes.NewReader(stripped), nil + case "jpg", "jpeg", "png": + if err := rejectOversizedImage(data); err != nil { + return nil, err } - return bytes.NewReader(data), nil - case "png": stripped, err := exifremove.Remove(data) if err == nil { return bytes.NewReader(stripped), nil @@ -40,6 +39,14 @@ func StripMetadata(file io.Reader, ext string) (*bytes.Reader, error) { } } +func rejectOversizedImage(data []byte) error { + err := validateImageDimensions(bytes.NewReader(data)) + if errors.Is(err, errImageDimensionsTooLarge) { + return fmt.Errorf("%w: %w", ErrInvalidImage, err) + } + return nil +} + func stripWEBPMetadata(data []byte) []byte { // Check if the file contains the RIFF...WEBP header if len(data) < 12 || string(data[:4]) != "RIFF" || string(data[8:12]) != "WEBP" { diff --git a/backend/internal/utils/image/profile_picture.go b/backend/internal/utils/image/profile_picture.go index c1435e88..82ecf878 100644 --- a/backend/internal/utils/image/profile_picture.go +++ b/backend/internal/utils/image/profile_picture.go @@ -24,6 +24,15 @@ var ErrInvalidImage = errors.New("invalid image") // CreateProfilePicture resizes the profile picture to a square and encodes it as PNG func CreateProfilePicture(file io.ReadSeeker) (io.ReadSeeker, error) { + // Reject an oversized pixel count before decoding can allocate a pixel buffer + validationErr := validateImageDimensions(file) + if _, err := file.Seek(0, io.SeekStart); err != nil { + return nil, fmt.Errorf("failed to seek file: %w", err) + } + if validationErr != nil { + return nil, fmt.Errorf("%w: %w", ErrInvalidImage, validationErr) + } + // Attempt standard formats first img, _, err := imageorient.Decode(file) if err != nil { diff --git a/backend/internal/utils/image/validation.go b/backend/internal/utils/image/validation.go new file mode 100644 index 00000000..79a42301 --- /dev/null +++ b/backend/internal/utils/image/validation.go @@ -0,0 +1,26 @@ +package profilepicture + +import ( + "errors" + "fmt" + "image" + "io" +) + +const maxImagePixels = 16_000_000 // e.g. 4000x4000 pixels + +var errImageDimensionsTooLarge = errors.New("image dimensions exceed the allowed limit") + +// validateImageDimensions checks if the image dimensions exceed the maximum allowed pixel count. +func validateImageDimensions(r io.Reader) error { + config, _, err := image.DecodeConfig(r) + if err != nil { + return err + } + + if int64(config.Width)*int64(config.Height) > maxImagePixels { + return fmt.Errorf("%w: got %dx%d, maximum pixel count is %d", errImageDimensionsTooLarge, config.Width, config.Height, maxImagePixels) + } + + return nil +} diff --git a/backend/internal/utils/image/validation_test.go b/backend/internal/utils/image/validation_test.go new file mode 100644 index 00000000..205e1a13 --- /dev/null +++ b/backend/internal/utils/image/validation_test.go @@ -0,0 +1,58 @@ +package profilepicture + +import ( + "bytes" + "encoding/binary" + "hash/crc32" + "image" + "image/png" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCreateProfilePictureRejectsOversizedPixelCountBeforeDecode(t *testing.T) { + _, err := CreateProfilePicture(bytes.NewReader(pngHeaderWithDimensions(12_000, 12_000))) + + require.ErrorIs(t, err, ErrInvalidImage) + require.ErrorIs(t, err, errImageDimensionsTooLarge) +} + +func TestCreateProfilePictureAcceptsImageWithinLimits(t *testing.T) { + var input bytes.Buffer + require.NoError(t, png.Encode(&input, image.NewNRGBA(image.Rect(0, 0, 400, 300)))) + + output, err := CreateProfilePicture(bytes.NewReader(input.Bytes())) + require.NoError(t, err) + + config, format, err := image.DecodeConfig(output) + require.NoError(t, err) + require.Equal(t, "png", format) + require.Equal(t, 300, config.Width) + require.Equal(t, 300, config.Height) +} + +func TestStripMetadataRejectsOversizedPixelCount(t *testing.T) { + _, err := StripMetadata(bytes.NewReader(pngHeaderWithDimensions(12_000, 12_000)), "png") + + require.ErrorIs(t, err, ErrInvalidImage) + require.ErrorIs(t, err, errImageDimensionsTooLarge) +} + +func pngHeaderWithDimensions(width, height uint32) []byte { + var out bytes.Buffer + out.Write([]byte("\x89PNG\r\n\x1a\n")) + + const headerLength = 13 + data := make([]byte, headerLength) + binary.BigEndian.PutUint32(data[0:4], width) + binary.BigEndian.PutUint32(data[4:8], height) + data[8] = 8 + + _ = binary.Write(&out, binary.BigEndian, uint32(headerLength)) + out.WriteString("IHDR") + out.Write(data) + _ = binary.Write(&out, binary.BigEndian, crc32.ChecksumIEEE(append([]byte("IHDR"), data...))) + + return out.Bytes() +}