Files
summercms/modules/cabana/crud_lifecycle_test.go
Jakub Zych 5e50b166ef refactor(10.2-01): nest framework packages under modules
- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
2026-09-28 02:21:02 +02:00

501 lines
17 KiB
Go

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
}
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")
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")
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")
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) 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
}