package web import ( "bytes" "context" "errors" "image" "image/color" "image/jpeg" "image/png" "io" "mime/multipart" "net/http" "net/http/httptest" "sync" "testing" "plumber/internal/blob" "plumber/internal/pacific" "plumber/internal/store" ) func TestPreparePostImage(t *testing.T) { t.Parallel() wide := solidPNG(t, 2000, 1000) prepared, err := preparePostImageHeader(t, "wide.png", wide) if err != nil { t.Fatal(err) } if prepared.width != 1600 || prepared.height != 800 || prepared.extension != ".png" || prepared.contentType != "image/png" { t.Fatalf("prepared PNG = %+v", prepared) } jpegBody := solidJPEG(t, 40, 20) prepared, err = preparePostImageHeader(t, "photo.jpg", jpegBody) if err != nil { t.Fatal(err) } if prepared.width != 40 || prepared.height != 20 || prepared.extension != ".jpg" || prepared.contentType != "image/jpeg" { t.Fatalf("prepared JPEG = %+v", prepared) } if _, err := preparePostImageHeader(t, "notes.txt", []byte("not an image")); err == nil { t.Fatal("text upload unexpectedly succeeded") } _, err = preparePostImageHeader(t, "too-large.jpg", make([]byte, postImageMaxFileBytes+1)) var requestErr *postImageRequestError if !errors.As(err, &requestErr) || requestErr.status != http.StatusRequestEntityTooLarge { t.Fatalf("oversized file error = %v, want 413 request error", err) } if _, err := preparePostImageHeader(t, "too-wide.png", solidPNG(t, 6001, 1)); err == nil { t.Fatal("oversized dimensions unexpectedly succeeded") } } func TestOrientPostImage(t *testing.T) { t.Parallel() source := image.NewNRGBA(image.Rect(0, 0, 2, 1)) source.Set(0, 0, color.NRGBA{R: 255, A: 255}) source.Set(1, 0, color.NRGBA{B: 255, A: 255}) rotated := orientPostImage(source, 6) if rotated.Bounds().Dx() != 1 || rotated.Bounds().Dy() != 2 { t.Fatalf("rotated bounds = %v", rotated.Bounds()) } top := color.NRGBAModel.Convert(rotated.At(0, 0)).(color.NRGBA) bottom := color.NRGBAModel.Convert(rotated.At(0, 1)).(color.NRGBA) if top.R != 255 || bottom.B != 255 { t.Fatalf("rotation colors top=%v bottom=%v", top, bottom) } } func TestPostImageMultipartLifecycle(t *testing.T) { t.Parallel() images := &recordingImageBlob{} srv, mem := newTestServer(t, Config{Blob: images}) handler := srv.Handler() homeowner := seedUser(t, mem, uniq("images"), "hunter22", store.RoleUser) cookies := loginUser(t, handler, homeowner.Username, "hunter22") csrf := csrfForCookies(t, handler, cookies) rec := multipartPost(t, handler, "/submit", map[string][]string{ "_csrf": {csrf}, "title": {"Leaky valve"}, "body": {"Two views of the leak."}, "city": {"Oakland"}, "image_description": {"Front view", "Under the sink"}, }, []multipartTestFile{ {name: "front.png", body: solidPNG(t, 80, 40)}, {name: "under.jpg", body: solidJPEG(t, 40, 80)}, }, cookies) if rec.Code != http.StatusSeeOther { t.Fatalf("root image upload status = %d: %s", rec.Code, rec.Body.String()) } roots, err := mem.ListRootPosts(context.Background(), pacific.Today(), homeowner.ID) if err != nil || len(roots) != 1 { t.Fatalf("roots = %+v, %v", roots, err) } root, err := mem.GetPost(context.Background(), roots[0].ID) if err != nil { t.Fatal(err) } if len(root.Images) != 2 || root.Images[0].Description != "Front view" || root.Images[1].Description != "Under the sink" { t.Fatalf("root images = %+v", root.Images) } rec = multipartPost(t, handler, "/posts", map[string][]string{ "_csrf": {csrf}, "parent_id": {root.ID}, "body": {"Here is the model label."}, "image_description": {"Model label"}, }, []multipartTestFile{{name: "label.png", body: solidPNG(t, 60, 30)}}, cookies) if rec.Code != http.StatusSeeOther { t.Fatalf("reply image upload status = %d: %s", rec.Code, rec.Body.String()) } thread, err := mem.GetPostThread(context.Background(), root.ID) if err != nil || len(thread.Replies) != 1 { t.Fatalf("thread = %+v, %v", thread, err) } reply := thread.Replies[0] if len(reply.Images) != 1 || reply.Images[0].Description != "Model label" { t.Fatalf("reply images = %+v", reply.Images) } originalKey := reply.Images[0].ObjectKey rec = multipartPost(t, handler, "/posts/"+reply.ID+"/edit", map[string][]string{ "_csrf": {csrf}, "body": {"Updated label photos."}, "existing_image_id": {reply.Images[0].ID}, "existing_image_description": {"Existing label"}, "image_description": {"Serial number"}, }, []multipartTestFile{{name: "serial.jpg", body: solidJPEG(t, 50, 25)}}, cookies) if rec.Code != http.StatusSeeOther { t.Fatalf("image edit status = %d: %s", rec.Code, rec.Body.String()) } edited, err := mem.GetPost(context.Background(), reply.ID) if err != nil { t.Fatal(err) } if len(edited.Images) != 2 || edited.Images[0].Description != "Existing label" || edited.Images[1].Description != "Serial number" { t.Fatalf("edited images = %+v", edited.Images) } rec = multipartPost(t, handler, "/posts/"+reply.ID+"/edit", map[string][]string{ "_csrf": {csrf}, "body": {"Keep only the serial number."}, "existing_image_id": {edited.Images[1].ID}, "existing_image_description": {"Serial number"}, }, nil, cookies) if rec.Code != http.StatusSeeOther { t.Fatalf("image removal status = %d: %s", rec.Code, rec.Body.String()) } edited, err = mem.GetPost(context.Background(), reply.ID) if err != nil { t.Fatal(err) } if len(edited.Images) != 1 || edited.Images[0].Description != "Serial number" { t.Fatalf("images after removal = %+v", edited.Images) } if !images.wasDeleted(originalKey) { t.Fatalf("removed object %q was not deleted: %+v", originalKey, images.deletedKeys()) } uploadsBefore := images.uploadCount() rec = multipartPost(t, handler, "/posts/"+reply.ID+"/edit", map[string][]string{ "_csrf": {csrf}, "body": {"Invalid retained image."}, "existing_image_id": {"not-owned"}, }, nil, cookies) if rec.Code != http.StatusBadRequest { t.Fatalf("invalid retained image status = %d, want 400", rec.Code) } unchanged, err := mem.GetPost(context.Background(), reply.ID) if err != nil { t.Fatal(err) } if unchanged.Body != "Keep only the serial number." || len(unchanged.Images) != 1 { t.Fatalf("invalid retained image changed post: %+v", unchanged) } files := make([]multipartTestFile, store.MaxPostImages+1) for i := range files { files[i] = multipartTestFile{name: "extra.png", body: solidPNG(t, 10, 10)} } rec = multipartPost(t, handler, "/posts", map[string][]string{ "_csrf": {csrf}, "parent_id": {root.ID}, "body": {"Too many images."}, }, files, cookies) if rec.Code != http.StatusBadRequest { t.Fatalf("five-image status = %d, want 400", rec.Code) } if images.uploadCount() != uploadsBefore { t.Fatal("five-image request uploaded objects before rejecting count") } } func TestPostImageUploadCompensation(t *testing.T) { t.Parallel() t.Run("blob failure deletes earlier upload", func(t *testing.T) { images := &recordingImageBlob{failAt: 2} srv, mem := newTestServer(t, Config{Blob: images}) handler := srv.Handler() user := seedUser(t, mem, uniq("blob-fail"), "hunter22", store.RoleUser) cookies := loginUser(t, handler, user.Username, "hunter22") csrf := csrfForCookies(t, handler, cookies) rec := multipartPost(t, handler, "/submit", map[string][]string{ "_csrf": {csrf}, "title": {"Upload failure"}, "body": {"Should not persist."}, }, []multipartTestFile{ {name: "one.png", body: solidPNG(t, 10, 10)}, {name: "two.png", body: solidPNG(t, 10, 10)}, }, cookies) if rec.Code != http.StatusServiceUnavailable { t.Fatalf("blob failure status = %d: %s", rec.Code, rec.Body.String()) } if images.uploadCount() != 1 || len(images.deletedKeys()) != 1 { t.Fatalf("blob compensation uploads=%d deletes=%v", images.uploadCount(), images.deletedKeys()) } }) t.Run("store failure deletes uploaded object", func(t *testing.T) { images := &recordingImageBlob{} mem := store.NewMemory() failing := &failingCreatePostStore{Store: mem} srv := newTestServerStore(t, failing, Config{Blob: images}) handler := srv.Handler() user := seedUser(t, mem, uniq("store-fail"), "hunter22", store.RoleUser) cookies := loginUser(t, handler, user.Username, "hunter22") csrf := csrfForCookies(t, handler, cookies) rec := multipartPost(t, handler, "/submit", map[string][]string{ "_csrf": {csrf}, "title": {"Store failure"}, "body": {"Should clean up."}, }, []multipartTestFile{{name: "one.png", body: solidPNG(t, 10, 10)}}, cookies) if rec.Code != http.StatusInternalServerError { t.Fatalf("store failure status = %d: %s", rec.Code, rec.Body.String()) } if images.uploadCount() != 1 || len(images.deletedKeys()) != 1 { t.Fatalf("store compensation uploads=%d deletes=%v", images.uploadCount(), images.deletedKeys()) } }) } func TestPostImageRequestLimits(t *testing.T) { t.Parallel() for _, test := range []struct { method string path string want int64 }{ {http.MethodPost, "/submit", postImageMaxRequestBytes}, {http.MethodPost, "/posts", postImageMaxRequestBytes}, {http.MethodPost, "/posts/id/edit", postImageMaxRequestBytes}, {http.MethodPost, "/login", defaultRequestBodyBytes}, {http.MethodGet, "/posts", defaultRequestBodyBytes}, } { req := httptest.NewRequest(test.method, test.path, nil) if got := requestBodyLimit(req); got != test.want { t.Errorf("%s %s limit = %d, want %d", test.method, test.path, got, test.want) } } } type multipartTestFile struct { name string body []byte } func multipartPost( t *testing.T, handler http.Handler, requestPath string, fields map[string][]string, files []multipartTestFile, cookies []*http.Cookie, ) *httptest.ResponseRecorder { t.Helper() var body bytes.Buffer writer := multipart.NewWriter(&body) for name, values := range fields { for _, value := range values { if err := writer.WriteField(name, value); err != nil { t.Fatal(err) } } } for _, file := range files { part, err := writer.CreateFormFile("images", file.name) if err != nil { t.Fatal(err) } if _, err := part.Write(file.body); err != nil { t.Fatal(err) } } if err := writer.Close(); err != nil { t.Fatal(err) } req := httptest.NewRequest(http.MethodPost, requestPath, &body) req.Header.Set("Content-Type", writer.FormDataContentType()) for _, cookie := range cookies { req.AddCookie(cookie) } rec := httptest.NewRecorder() handler.ServeHTTP(rec, req) return rec } func preparePostImageHeader(t *testing.T, name string, body []byte) (preparedPostImage, error) { t.Helper() var requestBody bytes.Buffer writer := multipart.NewWriter(&requestBody) part, err := writer.CreateFormFile("images", name) if err != nil { t.Fatal(err) } if _, err := part.Write(body); err != nil { t.Fatal(err) } if err := writer.Close(); err != nil { t.Fatal(err) } req := httptest.NewRequest(http.MethodPost, "/posts", &requestBody) req.Header.Set("Content-Type", writer.FormDataContentType()) if err := req.ParseMultipartForm(postImageMultipartMemory); err != nil { t.Fatal(err) } defer req.MultipartForm.RemoveAll() return preparePostImage(req.MultipartForm.File["images"][0]) } func solidPNG(t *testing.T, width, height int) []byte { t.Helper() img := image.NewNRGBA(image.Rect(0, 0, width, height)) for y := 0; y < height; y++ { for x := 0; x < width; x++ { img.Set(x, y, color.NRGBA{R: 30, G: 90, B: 140, A: 255}) } } var out bytes.Buffer if err := png.Encode(&out, img); err != nil { t.Fatal(err) } return out.Bytes() } func solidJPEG(t *testing.T, width, height int) []byte { t.Helper() img := image.NewNRGBA(image.Rect(0, 0, width, height)) for y := 0; y < height; y++ { for x := 0; x < width; x++ { img.Set(x, y, color.NRGBA{R: 140, G: 90, B: 30, A: 255}) } } var out bytes.Buffer if err := jpeg.Encode(&out, img, &jpeg.Options{Quality: 90}); err != nil { t.Fatal(err) } return out.Bytes() } type recordedImageUpload struct { key string contentType string body []byte } type recordingImageBlob struct { mu sync.Mutex calls int failAt int uploads []recordedImageUpload deletes []string } func (b *recordingImageBlob) Enabled() bool { return true } func (b *recordingImageBlob) Upload(_ context.Context, object blob.FileUpload) (string, error) { b.mu.Lock() defer b.mu.Unlock() b.calls++ if b.failAt > 0 && b.calls == b.failAt { return "", errors.New("injected upload failure") } body, err := io.ReadAll(object.Body) if err != nil { return "", err } b.uploads = append(b.uploads, recordedImageUpload{ key: object.Key, contentType: object.ContentType, body: body, }) return "https://cdn.example/" + object.Key, nil } func (b *recordingImageBlob) Delete(_ context.Context, key string) error { b.mu.Lock() defer b.mu.Unlock() b.deletes = append(b.deletes, key) return nil } func (b *recordingImageBlob) uploadCount() int { b.mu.Lock() defer b.mu.Unlock() return len(b.uploads) } func (b *recordingImageBlob) deletedKeys() []string { b.mu.Lock() defer b.mu.Unlock() return append([]string(nil), b.deletes...) } func (b *recordingImageBlob) wasDeleted(key string) bool { for _, deleted := range b.deletedKeys() { if deleted == key { return true } } return false } type failingCreatePostStore struct { store.Store } func (f *failingCreatePostStore) CreatePost(context.Context, *store.Post) error { return errors.New("injected store failure") }