package cabana import ( "bytes" "context" "fmt" "net/http" "net/http/httptest" "sort" "strings" "sync" "testing" "git.golem15.com/golem15/summercms/modules/bouncer" "gorm.io/gorm" ) func TestBulkDeleteEmpty(t *testing.T) { cap := &captureRouter{} (&service{}).mount(cap) key := "POST " + adminAPI("/{vendor}/{plugin}/{controller}/bulk-delete") mw, ok := cap.middleware[key] if !ok || !containsString(mw, "backend") { t.Fatalf("missing %s in %v", key, cap.routes) } _, httpSvc, _, db, hooks := hookFixture(t) hooks.perms = []string{"acme.demo.records"} forbidden := crudBulk(httpSvc, []byte(`{"ids":[1]}`), principalCtx(hooks, &bouncer.Principal{ID: 4})) if forbidden.Code != http.StatusForbidden || hooks.allocs.Load() != 0 || hooks.queries.Load() != 0 { t.Fatalf("forbidden=%d allocs=%d queries=%d body=%s", forbidden.Code, hooks.allocs.Load(), hooks.queries.Load(), forbidden.Body.String()) } ctx := principalCtx(hooks, superUser()) for _, body := range []string{`{"ids":[]}`, `{}`, `{"ids":null}`} { rec := crudBulk(httpSvc, []byte(body), ctx) if rec.Code != http.StatusUnprocessableEntity || !strings.Contains(rec.Body.String(), `"validation_failed"`) || !strings.Contains(rec.Body.String(), "The ids field is required.") { t.Fatalf("body %s => %d %s", body, rec.Code, rec.Body.String()) } } bad := crudBulk(httpSvc, []byte(`{"ids":["nope",1]}`), ctx) if bad.Code != http.StatusUnprocessableEntity || !strings.Contains(bad.Body.String(), `"validation_failed"`) { t.Fatalf("bad ids=%d %s", bad.Code, bad.Body.String()) } if n := countCrud(t, db); n != 0 { t.Fatalf("rows=%d", n) } } func TestBulkDeleteDuplicates(t *testing.T) { svc, _, cc, db, hooks := hookFixture(t) rows := seedCrud(t, db, 2) var got []uint ctx := idCtx(hooks, &got) res, err := svc.BulkDelete(ctx, cc, BulkDeleteInput{IDs: []any{rows[1].ID, rows[1].ID, rows[0].ID, float64(rows[0].ID)}}) if err != nil || res.Deleted != 2 { t.Fatalf("bulk=%+v err=%v", res, err) } if !sameUints(got, []uint{rows[0].ID, rows[1].ID}) { t.Fatalf("hook ids=%v, want each id once in ascending order", got) } if n := countCrud(t, db); n != 0 { t.Fatalf("rows=%d", n) } } func TestBulkDeleteOrder(t *testing.T) { svc, _, cc, db, hooks := hookFixture(t) rows := seedCrud(t, db, 3) var got []uint ctx := idCtx(hooks, &got) res, err := svc.BulkDelete(ctx, cc, BulkDeleteInput{IDs: []any{rows[2].ID, rows[0].ID, rows[1].ID}}) if err != nil || res.Deleted != 3 { t.Fatalf("bulk=%+v err=%v", res, err) } want := []uint{rows[0].ID, rows[1].ID, rows[2].ID} sort.Slice(want, func(i, j int) bool { return want[i] < want[j] }) if !sameUints(got, want) { t.Fatalf("hook ids=%v want %v", got, want) } } func TestBulkDeleteIdempotent(t *testing.T) { svc, _, cc, db, hooks := hookFixture(t) rows := seedCrud(t, db, 2) var got []uint ctx := idCtx(hooks, &got) first, err := svc.BulkDelete(ctx, cc, BulkDeleteInput{IDs: idsOf(rows)}) if err != nil || first.Deleted != 2 { t.Fatalf("first=%+v err=%v", first, err) } got = nil again, err := svc.BulkDelete(ctx, cc, BulkDeleteInput{IDs: idsOf(rows)}) if err != nil || again.Deleted != 0 || len(got) != 0 { t.Fatalf("repeat=%+v err=%v hooks=%v", again, err, got) } hooks.trackScope = true outside := seedCrud(t, db, 1) if err := db.Model(&outside[0]).Update("scope_id", 2).Error; err != nil { t.Fatal(err) } got = nil scoped := context.WithValue(ctx, scopeKey{}, uint(1)) res, err := svc.BulkDelete(scoped, cc, BulkDeleteInput{IDs: []any{outside[0].ID}}) if err != nil || res.Deleted != 0 || len(got) != 0 { t.Fatalf("out of scope=%+v err=%v hooks=%v", res, err, got) } if n := countCrud(t, db); n != 1 { t.Fatalf("rows=%d, out-of-scope delete changed the table", n) } } func TestBulkDeleteRollback(t *testing.T) { svc, httpSvc, cc, db, hooks := hookFixture(t) rows := seedCrud(t, db, 2) var got []uint ctx := idCtx(hooks, &got) if _, err := svc.BulkDelete(ctx, cc, BulkDeleteInput{IDs: []any{rows[0].ID, uint(999999)}}); err == nil || len(got) != 0 || countCrud(t, db) != 2 { t.Fatalf("mixed missing err=%v hooks=%v rows=%d", err, got, countCrud(t, db)) } rec := crudBulk(httpSvc, []byte(fmt.Sprintf(`{"ids":[%d,999999]}`, rows[0].ID)), ctx) if rec.Code != http.StatusConflict || strings.Contains(rec.Body.String(), "999999") { t.Fatalf("mixed http=%d %s", rec.Code, rec.Body.String()) } if countCrud(t, db) != 2 { t.Fatal("mixed http committed a delete") } hooks.trackScope = true if err := db.Model(&rows[1]).Update("scope_id", 2).Error; err != nil { t.Fatal(err) } got = nil scoped := context.WithValue(ctx, scopeKey{}, uint(1)) if _, err := svc.BulkDelete(scoped, cc, BulkDeleteInput{IDs: []any{rows[0].ID, rows[1].ID}}); err == nil || len(got) != 0 { t.Fatalf("mixed scope err=%v hooks=%v", err, got) } if loadCrud(t, db, rows[0].Name).ScopeID != 1 || loadCrud(t, db, rows[1].Name).ScopeID != 2 { t.Fatal("mixed scope committed") } if err := db.Model(&rows[1]).Update("scope_id", 1).Error; err != nil { t.Fatal(err) } hooks.trackScope = false got = nil failCtx := context.WithValue(ctx, failIDsKey{}, map[uint]bool{rows[1].ID: true}) _, err := svc.BulkDelete(failCtx, cc, BulkDeleteInput{IDs: idsOf(rows)}) if err == nil || strings.Contains(err.Error(), "secret-hook-boom") || !strings.Contains(err.Error(), "acme.demo.records") { t.Fatalf("hook err=%v", err) } if countCrud(t, db) != 2 { t.Fatal("hook failure committed a partial delete") } out := httptest.NewRecorder() writeCRUDError(out, err) if out.Code != http.StatusInternalServerError || strings.Contains(out.Body.String(), "secret-hook-boom") { t.Fatalf("hook http=%d %s", out.Code, out.Body.String()) } cancelCtx, cancel := context.WithCancel(ctx) cancelCtx = context.WithValue(cancelCtx, cancelKey{}, cancel) cancelCtx = context.WithValue(cancelCtx, cancelOnKey{}, rows[1].ID) got = nil if _, err := svc.BulkDelete(cancelCtx, cc, BulkDeleteInput{IDs: idsOf(rows)}); err == nil || countCrud(t, db) != 2 { t.Fatalf("cancel err=%v rows=%d", err, countCrud(t, db)) } } func TestBulkDeleteConcurrent(t *testing.T) { svc, _, cc, db, hooks := hookFixture(t) rows := seedCrud(t, db, 3) var got []uint ctx := context.WithValue(idCtx(hooks, &got), bulkSlowKey{}, true) in := BulkDeleteInput{IDs: idsOf(rows)} var wg sync.WaitGroup start := make(chan struct{}) type outcome struct { res BulkResult err error } outs := make([]outcome, 2) for i := 0; i < 2; i++ { wg.Add(1) go func(i int) { defer wg.Done() <-start outs[i].res, outs[i].err = svc.BulkDelete(ctx, cc, in) }(i) } close(start) wg.Wait() sum := 0 for i, out := range outs { if out.err != nil || (out.res.Deleted != 0 && out.res.Deleted != 3) { t.Fatalf("worker %d = %+v err=%v", i, out.res, out.err) } sum += out.res.Deleted } if sum != 3 || !sameUints(got, []uint{rows[0].ID, rows[1].ID, rows[2].ID}) || countCrud(t, db) != 0 { t.Fatalf("sum=%d hooks=%v rows=%d", sum, got, countCrud(t, db)) } } func seedCrud(t *testing.T, db *gorm.DB, n int) []crudRow { t.Helper() rows := make([]crudRow, n) for i := range rows { rows[i] = crudRow{Name: fmt.Sprintf("row-%d", i), ScopeID: 1} } if err := db.Create(&rows).Error; err != nil { t.Fatal(err) } return rows } func idsOf(rows []crudRow) []any { out := make([]any, len(rows)) for i := range rows { out[i] = rows[i].ID } return out } func idCtx(hooks *hookController, got *[]uint) context.Context { return context.WithValue(principalCtx(hooks, superUser()), idSinkKey{}, got) } func sameUints(got, want []uint) bool { if len(got) != len(want) { return false } for i := range got { if got[i] != want[i] { return false } } return true } func crudBulk(svc *service, body []byte, ctx context.Context) *httptest.ResponseRecorder { req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body)) req.SetPathValue("vendor", "acme") req.SetPathValue("plugin", "demo") req.SetPathValue("controller", "records") if ctx != nil { req = req.WithContext(ctx) } rec := httptest.NewRecorder() svc.bulkDelete(rec, req) return rec }