package cabana import ( "bytes" "context" "encoding/json" "net/http" "net/http/httptest" "strconv" "strings" "sync/atomic" "testing" "git.golem15.com/golem15/summercms/modules/backpack" "git.golem15.com/golem15/summercms/modules/bouncer" "git.golem15.com/golem15/summercms/modules/pact" "gorm.io/gorm" ) type scopeKey struct{} type hookController struct { log *[]string perms []string trackScope bool fail string queries atomic.Int32 allocs atomic.Int32 // txHooks counts the Before hooks that found the write's transaction on // their context through TxFromContext. txHooks atomic.Int32 } func (c *hookController) ID() string { return "acme.demo.records" } func (c *hookController) ModelName() string { return "Record" } func (c *hookController) ConfigDir() string { return "controllers/records" } func (c *hookController) RequiredPermissions() []string { if c == nil || c.perms == nil { return nil } return c.perms } func (c *hookController) NewRecord() any { if c != nil { c.allocs.Add(1) } return &crudRow{} } func (c *hookController) ListExtendQuery(ctx context.Context, db *gorm.DB) *gorm.DB { return c.FormExtendQuery(ctx, db) } func (c *hookController) FormExtendQuery(ctx context.Context, db *gorm.DB) *gorm.DB { if c != nil { c.queries.Add(1) } if c == nil || !c.trackScope || db == nil { return db } scope, _ := ctx.Value(scopeKey{}).(uint) return db.Where("scope_id = ?", scope) } func (c *hookController) FormBeforeCreate(ctx context.Context, model any) error { recordHook(ctx, "form_before_create") c.noteTx(ctx) if c.trackScope { if row, ok := model.(*crudRow); ok { if scope, ok := ctx.Value(scopeKey{}).(uint); ok { row.ScopeID = scope } } } return c.hookErr("form_before_create") } func (c *hookController) FormAfterCreate(ctx context.Context, model any) error { recordHook(ctx, "form_after_create") return c.hookErr("form_after_create") } func (c *hookController) FormBeforeUpdate(ctx context.Context, model any) error { recordHook(ctx, "form_before_update") c.noteTx(ctx) return c.hookErr("form_before_update") } func (c *hookController) FormAfterUpdate(ctx context.Context, model any) error { recordHook(ctx, "form_after_update") return c.hookErr("form_after_update") } func (c *hookController) FormBeforeDelete(ctx context.Context, model any) error { recordHook(ctx, "form_before_delete") c.noteTx(ctx) return c.hookErr("form_before_delete") } func (c *hookController) FormAfterDelete(ctx context.Context, model any) error { recordHook(ctx, "form_after_delete") return c.hookErr("form_after_delete") } func (c *hookController) noteTx(ctx context.Context) { if _, ok := TxFromContext(ctx); ok { c.txHooks.Add(1) } } func (c *hookController) hookErr(name string) error { if c != nil && c.fail == name { return errHookBoom } return nil } var errHookBoom = errString("secret-hook-boom") type errString string func (e errString) Error() string { return string(e) } func TestCRUDRecordRoutes(t *testing.T) { cap := &captureRouter{} (&service{}).mount(cap) for _, want := range []string{ "POST " + adminAPI("/{vendor}/{plugin}/{controller}"), "GET " + adminAPI("/{vendor}/{plugin}/{controller}/{id}"), "PUT " + adminAPI("/{vendor}/{plugin}/{controller}/{id}"), "DELETE " + adminAPI("/{vendor}/{plugin}/{controller}/{id}"), } { mw, ok := cap.middleware[want] if !ok { t.Fatalf("missing %s in %v", want, cap.routes) } if !containsString(mw, "backend") { t.Fatalf("%s middleware=%v, want backend before the handler", want, mw) } } _, httpSvc, _, db, hooks := hookFixture(t) ctx := principalCtx(hooks, superUser()) created := crudCall(httpSvc, http.MethodPost, "", []byte(`{"name":"Ada","note":"n"}`), ctx) if created.Code != http.StatusCreated { t.Fatalf("create status=%d body=%s", created.Code, created.Body.String()) } body := decodeData(t, created.Body.Bytes()) if body["name"] != "Ada" || body["note"] != "n" || body["id"] == nil { t.Fatalf("create data=%#v", body) } blank := crudCall(httpSvc, http.MethodPost, "", []byte(`{"name":""}`), ctx) if blank.Code != http.StatusUnprocessableEntity || !strings.Contains(blank.Body.String(), `"validation_failed"`) || !strings.Contains(blank.Body.String(), "The name field is required.") { t.Fatalf("blank create=%d %s", blank.Code, blank.Body.String()) } id := uintString(body["id"]) shown := crudCall(httpSvc, http.MethodGet, id, nil, ctx) if shown.Code != http.StatusOK || decodeData(t, shown.Body.Bytes())["name"] != "Ada" { t.Fatalf("show=%d %s", shown.Code, shown.Body.String()) } missing := crudCall(httpSvc, http.MethodGet, "999999", nil, ctx) if missing.Code != http.StatusNotFound || !strings.Contains(missing.Body.String(), `"not_found"`) { t.Fatalf("missing show=%d %s", missing.Code, missing.Body.String()) } bad := crudCall(httpSvc, http.MethodGet, "nope", nil, ctx) if bad.Code != http.StatusUnprocessableEntity || !strings.Contains(bad.Body.String(), `"validation_failed"`) { t.Fatalf("bad id=%d %s", bad.Code, bad.Body.String()) } updated := crudCall(httpSvc, http.MethodPut, id, []byte(`{"name":"Bea"}`), ctx) if updated.Code != http.StatusOK || decodeData(t, updated.Body.Bytes())["name"] != "Bea" { t.Fatalf("update=%d %s", updated.Code, updated.Body.String()) } *hooks.log = nil deleted := crudCall(httpSvc, http.MethodDelete, id, nil, ctx) if deleted.Code != http.StatusOK || deletedCount(t, deleted.Body.Bytes()) != 1 { t.Fatalf("delete=%d %s", deleted.Code, deleted.Body.String()) } if strings.Join(*hooks.log, ",") != "form_before_delete,before_delete,after_delete,form_after_delete" { t.Fatalf("delete hooks=%v", *hooks.log) } *hooks.log = nil again := crudCall(httpSvc, http.MethodDelete, id, nil, ctx) if again.Code != http.StatusOK || deletedCount(t, again.Body.Bytes()) != 0 || len(*hooks.log) != 0 { t.Fatalf("repeat delete=%d %s hooks=%v", again.Code, again.Body.String(), *hooks.log) } if n := countCrud(t, db); n != 0 { t.Fatalf("rows=%d after delete", n) } } func TestCRUDPermissions(t *testing.T) { _, httpSvc, _, db, hooks := hookFixture(t) hooks.perms = []string{"acme.demo.records"} denied := &bouncer.Principal{ID: 4, PermissionGrants: map[string]bool{"other.code": true}} unauth := crudCall(httpSvc, http.MethodPost, "", []byte(`{`), principalCtx(hooks, nil)) if unauth.Code != http.StatusUnauthorized || !strings.Contains(unauth.Body.String(), `"unauthenticated"`) { t.Fatalf("unauth=%d %s", unauth.Code, unauth.Body.String()) } forbidden := crudCall(httpSvc, http.MethodPost, "", []byte(`{`), principalCtx(hooks, denied)) if forbidden.Code != http.StatusForbidden || !strings.Contains(forbidden.Body.String(), `"forbidden"`) { t.Fatalf("forbidden=%d %s", forbidden.Code, forbidden.Body.String()) } badID := crudCall(httpSvc, http.MethodGet, "nope", nil, principalCtx(hooks, denied)) if badID.Code != http.StatusForbidden { t.Fatalf("bad id before permission=%d %s", badID.Code, badID.Body.String()) } if hooks.allocs.Load() != 0 || hooks.queries.Load() != 0 { t.Fatalf("allocs=%d queries=%d, permission did not run first", hooks.allocs.Load(), hooks.queries.Load()) } if n := countCrud(t, db); n != 0 { t.Fatalf("denied request persisted %d rows", n) } ok := crudCall(httpSvc, http.MethodPost, "", []byte(`{"name":"Ada"}`), principalCtx(hooks, superUser())) if ok.Code != http.StatusCreated { t.Fatalf("superuser create=%d %s", ok.Code, ok.Body.String()) } } func TestCRUDScope(t *testing.T) { _, httpSvc, _, db, hooks := hookFixture(t) hooks.trackScope = true in := &crudRow{Name: "In", ScopeID: 1} out := &crudRow{Name: "Out", ScopeID: 2} if err := db.Create(in).Error; err != nil { t.Fatal(err) } if err := db.Create(out).Error; err != nil { t.Fatal(err) } ctx := principalCtx(hooks, superUser()) ctx = context.WithValue(ctx, scopeKey{}, uint(1)) shown := crudCall(httpSvc, http.MethodGet, idString(in.ID), nil, ctx) if shown.Code != http.StatusOK || decodeData(t, shown.Body.Bytes())["name"] != "In" { t.Fatalf("in scope show=%d %s", shown.Code, shown.Body.String()) } if _, ok := decodeData(t, shown.Body.Bytes())["scope_id"]; ok { t.Fatalf("show leaked scope_id: %s", shown.Body.String()) } missing := crudCall(httpSvc, http.MethodGet, "999999", nil, ctx) hidden := crudCall(httpSvc, http.MethodGet, idString(out.ID), nil, ctx) if missing.Code != http.StatusNotFound || hidden.Code != http.StatusNotFound || missing.Body.String() != hidden.Body.String() { t.Fatalf("missing=%d %s hidden=%d %s", missing.Code, missing.Body.String(), hidden.Code, hidden.Body.String()) } hacked := crudCall(httpSvc, http.MethodPut, idString(out.ID), []byte(`{"name":"hacked"}`), ctx) if hacked.Code != http.StatusNotFound || hacked.Body.String() != missing.Body.String() { t.Fatalf("out of scope update=%d %s", hacked.Code, hacked.Body.String()) } if got := loadCrud(t, db, "Out"); got.Name != "Out" || got.ScopeID != 2 { t.Fatalf("out of scope row changed: %+v", got) } *hooks.log = nil gone := crudCall(httpSvc, http.MethodDelete, "999999", nil, ctx) hiddenDelete := crudCall(httpSvc, http.MethodDelete, idString(out.ID), nil, ctx) if gone.Code != http.StatusOK || hiddenDelete.Code != http.StatusOK || deletedCount(t, gone.Body.Bytes()) != 0 || deletedCount(t, hiddenDelete.Body.Bytes()) != 0 || gone.Body.String() != hiddenDelete.Body.String() { t.Fatalf("delete missing=%d %s hidden=%d %s", gone.Code, gone.Body.String(), hiddenDelete.Code, hiddenDelete.Body.String()) } if len(*hooks.log) != 0 { t.Fatalf("delete hooks ran for absent rows: %v", *hooks.log) } if got := loadCrud(t, db, "Out"); got.ScopeID != 2 { t.Fatalf("out of scope delete removed %+v", got) } } func TestCRUDHooks(t *testing.T) { svc, _, cc, _, hooks := hookFixture(t) ctx := principalCtx(hooks, superUser()) rec, err := svc.Create(ctx, cc, RecordInput{Body: map[string]any{"name": "Ada", "note": "n"}}) if err != nil { t.Fatalf("create err=%v", err) } if got := strings.Join(*hooks.log, ","); got != "before_validate,form_before_create,before_save,before_create,after_create,after_save,form_after_create" { t.Fatalf("create hooks=%s", got) } if countHooks(*hooks.log, "form_before_create") != 1 || countHooks(*hooks.log, "before_create") != 1 || countHooks(*hooks.log, "form_after_create") != 1 { t.Fatalf("create hook repeated: %v", *hooks.log) } *hooks.log = nil if _, err := svc.Update(ctx, cc, rec["id"], RecordInput{Body: map[string]any{"name": "Bea"}}); err != nil { t.Fatalf("update err=%v", err) } if got := strings.Join(*hooks.log, ","); got != "before_validate,form_before_update,before_save,before_update,after_update,after_save,form_after_update" { t.Fatalf("update hooks=%s", got) } *hooks.log = nil if _, err := svc.Delete(ctx, cc, rec["id"]); err != nil { t.Fatalf("delete err=%v", err) } if got := strings.Join(*hooks.log, ","); got != "form_before_delete,before_delete,after_delete,form_after_delete" { t.Fatalf("delete hooks=%s", got) } } func TestCRUDRollback(t *testing.T) { svc, httpSvc, cc, db, hooks := hookFixture(t) ctx := principalCtx(hooks, superUser()) hooks.fail = "form_after_create" _, err := svc.Create(ctx, cc, RecordInput{Body: map[string]any{"name": "Ada"}}) if err == nil { t.Fatal("after-create hook failure was ignored") } if strings.Contains(err.Error(), "secret-hook-boom") || !strings.Contains(err.Error(), "acme.demo.records") { t.Fatalf("hook error=%v, want opaque controller context", err) } rec := httptest.NewRecorder() writeCRUDError(rec, err) if rec.Code != http.StatusInternalServerError || strings.Contains(rec.Body.String(), "secret-hook-boom") || strings.Contains(rec.Body.String(), "acme.demo.records") { t.Fatalf("http hook error=%d %s", rec.Code, rec.Body.String()) } if n := countCrud(t, db); n != 0 { t.Fatalf("rows=%d after failed create", n) } hooks.fail = "" created, err := svc.Create(ctx, cc, RecordInput{Body: map[string]any{"name": "Ada"}}) if err != nil { t.Fatal(err) } hooks.fail = "form_before_update" if _, err := svc.Update(ctx, cc, created["id"], RecordInput{Body: map[string]any{"name": "hacked"}}); err == nil || strings.Contains(err.Error(), "secret-hook-boom") { t.Fatalf("update hook err=%v", err) } if got := loadCrud(t, db, "Ada"); got.Name != "Ada" { t.Fatalf("failed update committed %+v", got) } hooks.fail = "form_after_update" if _, err := svc.Update(ctx, cc, created["id"], RecordInput{Body: map[string]any{"note": "later"}}); err == nil { t.Fatal("after-update hook failure was ignored") } if got := loadCrud(t, db, "Ada"); got.Note != "n" && got.Note != "" { t.Fatalf("failed after-update committed note %q", got.Note) } if got := loadCrud(t, db, "Ada"); got.Note == "later" { t.Fatal("after-update wrote note") } *hooks.log = nil failCtx := context.WithValue(ctx, failDeleteKey{}, true) if _, err := svc.Delete(failCtx, cc, created["id"]); err == nil || strings.Contains(err.Error(), "secret-hook-boom") { t.Fatalf("delete hook err=%v", err) } if !containsString(*hooks.log, "before_delete") { t.Fatalf("delete lifecycle did not run: %v", *hooks.log) } if got := loadCrud(t, db, "Ada"); got.Name != "Ada" { t.Fatalf("failed delete removed %+v", got) } _ = httpSvc } type captureRouter struct { prefix string mw []string routes []string middleware map[string][]string } func (c *captureRouter) Group(prefix string, middleware []string, fn func(pact.Router)) { c.GroupRaw(prefix, middleware, fn) } func (c *captureRouter) GroupRaw(prefix string, middleware []string, fn func(pact.Router)) { if c.middleware == nil { c.middleware = map[string][]string{} } child := &captureRouter{ prefix: c.prefix + prefix, mw: append(append([]string{}, c.mw...), middleware...), middleware: c.middleware, } fn(child) c.routes = append(c.routes, child.routes...) } func (c *captureRouter) Get(path string, _ http.HandlerFunc, _ ...string) { c.add("GET", path) } func (c *captureRouter) Post(path string, _ http.HandlerFunc, _ ...string) { c.add("POST", path) } func (c *captureRouter) Put(path string, _ http.HandlerFunc, _ ...string) { c.add("PUT", path) } func (c *captureRouter) Patch(path string, _ http.HandlerFunc, _ ...string) { c.add("PATCH", path) } func (c *captureRouter) Delete(path string, _ http.HandlerFunc, _ ...string) { c.add("DELETE", path) } func (c *captureRouter) Where(string, string) {} func (c *captureRouter) WhereIn(string, ...string) {} func (c *captureRouter) add(method, path string) { full := method + " " + c.prefix + path c.routes = append(c.routes, full) if c.middleware == nil { c.middleware = map[string][]string{} } c.middleware[full] = append([]string{}, c.mw...) } func hookFixture(t *testing.T) (CRUDService, *service, *CompiledController, *gorm.DB, *hookController) { t.Helper() crudSvc, cc, db := crudFixture(t) hooks := &hookController{log: &[]string{}} cc.Controller = hooks app := backpack.New(nil) if err := app.Publish(db); err != nil { t.Fatal(err) } httpSvc := &service{app: app, reg: &Registry{byID: map[string]*CompiledController{ cc.Controller.ID(): cc, }}} return crudSvc, httpSvc, cc, db, hooks } func principalCtx(hooks *hookController, principal *bouncer.Principal) context.Context { ctx := context.WithValue(context.Background(), hookSinkKey{}, hooks.log) if principal != nil { principal.Backend = true ctx = bouncer.WithUser(ctx, principal) } return ctx } func superUser() *bouncer.Principal { return &bouncer.Principal{ID: 1, Backend: true, IsSuperuser: true} } func crudCall(svc *service, method, id string, body []byte, ctx context.Context) *httptest.ResponseRecorder { req := httptest.NewRequest(method, "/", bytes.NewReader(body)) req.SetPathValue("vendor", "acme") req.SetPathValue("plugin", "demo") req.SetPathValue("controller", "records") if id != "" { req.SetPathValue("id", id) } if ctx != nil { req = req.WithContext(ctx) } rec := httptest.NewRecorder() switch method { case http.MethodPost: svc.create(rec, req) case http.MethodGet: svc.show(rec, req) case http.MethodPut: svc.update(rec, req) case http.MethodDelete: svc.deleteRecord(rec, req) default: rec.Code = http.StatusNotImplemented } return rec } func decodeData(t *testing.T, raw []byte) map[string]any { t.Helper() var body struct { Data map[string]any `json:"data"` Meta map[string]any `json:"meta"` } if err := jsonUnmarshal(raw, &body); err != nil { t.Fatalf("json: %v body=%s", err, raw) } if body.Meta == nil { t.Fatalf("meta is null: %s", raw) } return body.Data } func deletedCount(t *testing.T, raw []byte) float64 { t.Helper() var body struct { Data struct { Deleted float64 `json:"deleted"` } `json:"data"` } if err := jsonUnmarshal(raw, &body); err != nil { t.Fatalf("json: %v body=%s", err, raw) } return body.Data.Deleted } func jsonUnmarshal(raw []byte, dest any) error { return json.Unmarshal(raw, dest) } func uintString(id any) string { switch n := id.(type) { case uint: return idString(n) case float64: return idString(uint(n)) case int: return idString(uint(n)) default: return "" } } func idString(id uint) string { return strconv.FormatUint(uint64(id), 10) } func countHooks(log []string, name string) int { n := 0 for _, item := range log { if item == name { n++ } } return n } func containsString(items []string, want string) bool { for _, item := range items { if item == want { return true } } return false } // TestCRUDOperationsFollowDeclarations pins WR-03: the compiled list and form // decide which writes the server accepts, not only what the SPA shows. func TestCRUDOperationsFollowDeclarations(t *testing.T) { _, httpSvc, cc, db, hooks := hookFixture(t) ctx := principalCtx(hooks, superUser()) created := crudCall(httpSvc, http.MethodPost, "", []byte(`{"name":"Ada"}`), ctx) if created.Code != http.StatusCreated { t.Fatalf("declared create=%d %s", created.Code, created.Body.String()) } id := uintString(decodeData(t, created.Body.Bytes())["id"]) bulk := func() *httptest.ResponseRecorder { req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"ids":[`+id+`]}`)).WithContext(ctx) req.SetPathValue("vendor", "acme") req.SetPathValue("plugin", "demo") req.SetPathValue("controller", "records") rec := httptest.NewRecorder() httpSvc.bulkDelete(rec, req) return rec } forbidden := func(name string, rec *httptest.ResponseRecorder) { t.Helper() if rec.Code != http.StatusForbidden || !strings.Contains(rec.Body.String(), `"forbidden"`) { t.Fatalf("%s=%d %s, want 403 forbidden", name, rec.Code, rec.Body.String()) } } // No toolbar create button: the create route is closed, update is not. list := *cc.List cc.List = &list cc.List.ToolbarButtons = []string{"delete"} forbidden("create without a toolbar create button", crudCall(httpSvc, http.MethodPost, "", []byte(`{"name":"Bea"}`), ctx)) if got := crudCall(httpSvc, http.MethodPut, id, []byte(`{"name":"Cid"}`), ctx); got.Code != http.StatusOK { t.Fatalf("update=%d %s", got.Code, got.Body.String()) } // No toolbar delete button: bulk delete is closed. cc.List.ToolbarButtons = []string{"create"} forbidden("bulk delete without a toolbar delete button", bulk()) if n := countCrud(t, db); n != 1 { t.Fatalf("rows=%d, a refused bulk delete removed data", n) } // No form: nothing can be created, updated or deleted one by one. form := cc.Form cc.Form = nil forbidden("create without a form", crudCall(httpSvc, http.MethodPost, "", []byte(`{"name":"Dan"}`), ctx)) forbidden("update without a form", crudCall(httpSvc, http.MethodPut, id, []byte(`{"name":"Eve"}`), ctx)) forbidden("delete without a form", crudCall(httpSvc, http.MethodDelete, id, nil, ctx)) if n := countCrud(t, db); n != 1 { t.Fatalf("rows=%d, a refused write changed data", n) } cc.Form = form // Declared again: bulk delete works. cc.List.ToolbarButtons = []string{"create", "delete"} if got := bulk(); got.Code != http.StatusOK || deletedCount(t, got.Body.Bytes()) != 1 { t.Fatalf("declared bulk delete=%d %s", got.Code, got.Body.String()) } } // TestHooksReceiveTheWriteTransaction pins WR-19: lifecycle hooks can reach the // transaction their write runs in through TxFromContext, so a hook's own reads // share the write's snapshot and connection instead of using the app pool. func TestHooksReceiveTheWriteTransaction(t *testing.T) { _, httpSvc, _, _, hooks := hookFixture(t) ctx := principalCtx(hooks, superUser()) created := crudCall(httpSvc, http.MethodPost, "", []byte(`{"name":"Ada"}`), ctx) if created.Code != http.StatusCreated { t.Fatalf("create=%d %s", created.Code, created.Body.String()) } id := uintString(decodeData(t, created.Body.Bytes())["id"]) if updated := crudCall(httpSvc, http.MethodPut, id, []byte(`{"name":"Bea"}`), ctx); updated.Code != http.StatusOK { t.Fatalf("update=%d %s", updated.Code, updated.Body.String()) } if deleted := crudCall(httpSvc, http.MethodDelete, id, nil, ctx); deleted.Code != http.StatusOK { t.Fatalf("delete=%d %s", deleted.Code, deleted.Body.String()) } if got := hooks.txHooks.Load(); got != 3 { t.Fatalf("Before hooks that saw the transaction = %d, want 3 (create, update, delete)", got) } if _, ok := TxFromContext(context.Background()); ok { t.Fatal("a bare context reported a transaction") } if _, ok := TxFromContext(nil); ok { //nolint:staticcheck // a nil context must not panic t.Fatal("a nil context reported a transaction") } }