diff --git a/cabana/crud.go b/cabana/crud.go index e9c1163..139863f 100644 --- a/cabana/crud.go +++ b/cabana/crud.go @@ -2,10 +2,18 @@ package cabana import ( "context" + "encoding/json" "errors" + "math" "net/http" + "reflect" + "strconv" + "strings" + "git.golem15.com/golem15/summercms/lagoon" + "git.golem15.com/golem15/summercms/pact" "gorm.io/gorm" + "gorm.io/gorm/clause" ) // CRUDService runs schema-projected record and bulk operations. @@ -41,31 +49,481 @@ type CapabilityError struct { } func (e *CapabilityError) Error() string { + if e == nil { + return "cabana: missing Fill/Validate capability" + } return "cabana: controller " + e.ControllerID + ": missing Fill/Validate capability" } -var errCRUDNotReady = errors.New("cabana: crud is not implemented") +type recordNotFound struct{} -// ProjectWritableFields copies only activation-bound writable keys. -func ProjectWritableFields(cc *CompiledController, body map[string]any) map[string]any { - return nil +func (recordNotFound) Error() string { return "cabana: not found" } + +type hasRules interface { + Rules() map[string]string } -// BindWritableFields records schema field names onto model fill keys. +// ProjectWritableFields copies only activation-bound writable keys. +// Unknown keys, case variants, nested objects, and protected columns are dropped. +func ProjectWritableFields(cc *CompiledController, body map[string]any) map[string]any { + out := map[string]any{} + if cc == nil || body == nil { + return out + } + allowed := map[string]string{} + for _, field := range cc.Writable { + if field.FillKey == "" || protectedFillKey(field.Name) || protectedFillKey(field.FillKey) { + continue + } + allowed[field.Name] = field.FillKey + } + for key, val := range body { + fillKey, ok := allowed[key] + if !ok || nestedValue(val) { + continue + } + out[fillKey] = val + } + return out +} + +// BindWritableFields records schema field names onto model column fill keys. +// Protected columns are omitted. A scalar field with no column fails activation. func BindWritableFields(cc *CompiledController) error { - return errCRUDNotReady + if cc == nil || cc.Form == nil { + return nil + } + id := "" + if cc.Controller != nil { + id = cc.Controller.ID() + } + src, ok := cc.Controller.(pact.AdminRecordSource) + if !ok || src == nil || src.NewRecord() == nil { + return nil + } + cols := modelColumns(src.NewRecord()) + bindings := make([]WritableField, 0, len(cc.Form.Fields)) + for _, field := range cc.Form.Fields { + if !scalarFormField(field.Type) || protectedFillKey(field.Name) { + continue + } + if _, known := cols[field.Name]; !known { + return errors.New("cabana: controller " + id + ": field " + field.Name + " is not a model column") + } + bindings = append(bindings, WritableField{Name: field.Name, FillKey: field.Name}) + } + cc.Writable = bindings + return nil } // Create persists a projected record after Fill and Validate. func (s CRUDService) Create(ctx context.Context, cc *CompiledController, in RecordInput) (map[string]any, error) { - return nil, errCRUDNotReady + return s.save(ctx, cc, nil, in, false) } // Update persists a projected change after Fill and Validate. func (s CRUDService) Update(ctx context.Context, cc *CompiledController, id any, in RecordInput) (map[string]any, error) { - return nil, errCRUDNotReady + return s.save(ctx, cc, id, in, true) +} + +func (s CRUDService) save(ctx context.Context, cc *CompiledController, id any, in RecordInput, update bool) (map[string]any, error) { + if s.DB == nil { + return nil, errors.New("cabana: database is not configured") + } + if ctx == nil { + ctx = context.Background() + } + if _, err := newWritableModel(cc); err != nil { + return nil, err + } + var result map[string]any + err := s.DB.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + tx = tx.WithContext(ctx) + target, err := newWritableModel(cc) + if err != nil { + return err + } + op := "create" + if update { + op = "update" + pk, err := coercePK(target, id) + if err != nil { + return err + } + if err := loadRecord(ctx, tx, cc, target, pk); err != nil { + return err + } + } + projected := projectOperation(cc, in.Body, op) + if err := lagoon.Fill(target, fillAllowed(cc, target, op), projected, false); err != nil { + return &CapabilityError{ControllerID: controllerID(cc)} + } + if hook, ok := target.(lagoon.HasBeforeValidate); ok && hook != nil { + if err := hook.BeforeValidate(tx); err != nil { + return &CapabilityError{ControllerID: controllerID(cc)} + } + } + rules := mergedRules(cc, target) + msgs, err := lagoon.Validate(ctx, tx, target, rules, valuesForRules(target, rules), nil) + if err != nil { + return &CapabilityError{ControllerID: controllerID(cc)} + } + if len(msgs) > 0 { + return &ValidationError{Details: validationDetails(msgs)} + } + if update { + err = tx.Save(target).Error + } else { + err = tx.Create(target).Error + } + if err != nil { + return err + } + result = projectRecord(cc, target) + return nil + }) + if err != nil { + return nil, err + } + return result, nil } func writeCRUDError(w http.ResponseWriter, err error) { + var ve *ValidationError + if errors.As(err, &ve) { + WriteErrorDetails(w, http.StatusUnprocessableEntity, "validation_failed", "Validation failed", ve.Details) + return + } + var missing recordNotFound + if errors.As(err, &missing) { + WriteError(w, http.StatusNotFound, "not_found", msgNotFound) + return + } WriteError(w, http.StatusInternalServerError, "error", msgServerError) } + +func newWritableModel(cc *CompiledController) (any, error) { + id := controllerID(cc) + if cc == nil || cc.Controller == nil { + return nil, &CapabilityError{ControllerID: id} + } + src, ok := cc.Controller.(pact.AdminRecordSource) + if !ok || src == nil { + return nil, &CapabilityError{ControllerID: id} + } + model := src.NewRecord() + if model == nil { + return nil, &CapabilityError{ControllerID: id} + } + if _, ok := model.(lagoon.HasFillable); !ok { + return nil, &CapabilityError{ControllerID: id} + } + if _, ok := model.(hasRules); !ok { + return nil, &CapabilityError{ControllerID: id} + } + return model, nil +} + +func controllerID(cc *CompiledController) string { + if cc == nil || cc.Controller == nil { + return "" + } + return cc.Controller.ID() +} + +func loadRecord(ctx context.Context, tx *gorm.DB, cc *CompiledController, dest any, pk any) error { + col := primaryColumn(dest) + q := tx.WithContext(ctx).Clauses(clause.Locking{Strength: "UPDATE"}) + err := q.Where(clause.Eq{Column: clause.Column{Name: col}, Value: pk}).Take(dest).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return recordNotFound{} + } + return err +} + +func projectOperation(cc *CompiledController, body map[string]any, op string) map[string]any { + projected := ProjectWritableFields(cc, body) + if op == "" || cc == nil { + return projected + } + filtered := map[string]any{} + for _, field := range cc.Writable { + if !contextAllows(cc, field.Name, op) { + continue + } + if val, ok := projected[field.FillKey]; ok { + filtered[field.FillKey] = val + } + } + return filtered +} + +func fillAllowed(cc *CompiledController, model any, op string) []string { + fillable, ok := model.(lagoon.HasFillable) + if !ok || fillable == nil { + return nil + } + allowed := map[string]struct{}{} + for _, key := range fillable.Fillable() { + if !protectedFillKey(key) { + allowed[key] = struct{}{} + } + } + out := make([]string, 0, len(allowed)) + if cc == nil { + return out + } + for _, field := range cc.Writable { + if protectedFillKey(field.FillKey) { + continue + } + if _, ok := allowed[field.FillKey]; !ok || !contextAllows(cc, field.Name, op) { + continue + } + out = append(out, field.FillKey) + } + return out +} + +func mergedRules(cc *CompiledController, model any) map[string]string { + out := map[string]string{} + if rules, ok := model.(hasRules); ok && rules != nil { + for key, rule := range rules.Rules() { + out[key] = rule + } + } + if cc == nil || cc.Form == nil { + return out + } + for _, field := range cc.Form.Fields { + if field.Required { + out[field.Name] = mergeRequired(out[field.Name]) + } + } + return out +} + +func mergeRequired(rule string) string { + if strings.TrimSpace(rule) == "" { + return "required" + } + for _, tok := range strings.Split(rule, "|") { + name, _, _ := strings.Cut(strings.TrimSpace(tok), ":") + if name == "required" { + return rule + } + } + return rule + "|required" +} + +func valuesForRules(model any, rules map[string]string) map[string]any { + v := reflect.ValueOf(model) + for v.Kind() == reflect.Pointer { + if v.IsNil() { + return map[string]any{} + } + v = v.Elem() + } + out := map[string]any{} + for field := range rules { + f := fieldByColumn(v, field) + if f.IsValid() && f.CanInterface() { + out[field] = f.Interface() + } + } + return out +} + +func validationDetails(msgs map[string][]string) map[string]any { + out := make(map[string]any, len(msgs)) + for key, messages := range msgs { + out[key] = messages + } + return out +} + +func projectRecord(cc *CompiledController, model any) map[string]any { + v := reflect.ValueOf(model) + for v.Kind() == reflect.Pointer { + if v.IsNil() { + return map[string]any{} + } + v = v.Elem() + } + out := map[string]any{} + if id := fieldByColumn(v, primaryColumn(model)); id.IsValid() && id.CanInterface() { + out["id"] = id.Interface() + } + if cc == nil { + return out + } + for _, field := range cc.Writable { + if protectedFillKey(field.FillKey) { + continue + } + f := fieldByColumn(v, field.FillKey) + if f.IsValid() && f.CanInterface() { + out[field.Name] = f.Interface() + } + } + return out +} + +func modelColumns(model any) map[string]struct{} { + t := reflect.TypeOf(model) + for t != nil && t.Kind() == reflect.Pointer { + t = t.Elem() + } + cols := map[string]struct{}{} + if t == nil || t.Kind() != reflect.Struct { + return cols + } + for i := 0; i < t.NumField(); i++ { + field := t.Field(i) + if field.PkgPath != "" { + continue + } + name := gormColumn(field) + if name != "" { + cols[name] = struct{}{} + } + } + return cols +} + +func scalarFormField(typ string) bool { + switch typ { + case "text", "textarea", "number", "checkbox", "switch", "dropdown": + return true + default: + return false + } +} + +func protectedFillKey(key string) bool { + switch strings.ToLower(key) { + case "id", "created_at", "updated_at", "deleted_at", + "owner_id", "user_id", "collection_id", "organisation_id", "organization_id", + "scope_id", "role_id", "permissions", + "is_superuser", "is_system", "is_activated", "password": + return true + default: + return false + } +} + +func nestedValue(val any) bool { + switch val.(type) { + case map[string]any, []any: + return true + default: + return false + } +} + +func contextAllows(cc *CompiledController, name, op string) bool { + if op == "" || cc == nil || cc.Form == nil { + return true + } + for _, field := range cc.Form.Fields { + if field.Name != name { + continue + } + if field.Context == nil || len(field.Context.values) == 0 { + return true + } + for _, value := range field.Context.values { + if value == op { + return true + } + } + return false + } + return true +} + +func coercePK(model any, id any) (any, error) { + n, err := asUint(id) + if err != nil { + return nil, &ValidationError{Details: map[string]any{"id": []string{"The id field must be an integer."}}} + } + return castPK(model, n), nil +} + +func castPK(model any, n uint) any { + t := reflect.TypeOf(model) + for t != nil && t.Kind() == reflect.Pointer { + t = t.Elem() + } + if t == nil || t.Kind() != reflect.Struct { + return n + } + for i := 0; i < t.NumField(); i++ { + field := t.Field(i) + if !strings.Contains(field.Tag.Get("gorm"), "primaryKey") { + continue + } + switch field.Type.Kind() { + case reflect.Uint: + return uint(n) + case reflect.Uint32: + return uint32(n) + case reflect.Uint64: + return uint64(n) + case reflect.Int: + return int(n) + case reflect.Int64: + return int64(n) + case reflect.String: + return strconv.FormatUint(uint64(n), 10) + default: + return n + } + } + return n +} + +func asUint(id any) (uint, error) { + switch n := id.(type) { + case uint: + return n, nil + case uint32: + return uint(n), nil + case uint64: + if uint64(uint(n)) != n { + return 0, errBadID + } + return uint(n), nil + case int: + if n < 0 { + return 0, errBadID + } + return uint(n), nil + case int64: + if n < 0 { + return 0, errBadID + } + return uint(n), nil + case float64: + if n < 0 || n != math.Trunc(n) || n > math.MaxUint32 && strconv.IntSize == 32 { + return 0, errBadID + } + return uint(n), nil + case json.Number: + i, err := n.Int64() + if err != nil || i < 0 { + return 0, errBadID + } + return uint(i), nil + case string: + i, err := strconv.ParseUint(n, 10, 64) + if err != nil { + return 0, errBadID + } + return uint(i), nil + default: + return 0, errBadID + } +} + +var errBadID = errors.New("cabana: bad id") diff --git a/cabana/registry.go b/cabana/registry.go index d98e242..54275db 100644 --- a/cabana/registry.go +++ b/cabana/registry.go @@ -53,12 +53,16 @@ func compileRegistry(items []controllerRef) (*Registry, error) { if err != nil { return nil, err } - byID[id] = &CompiledController{ + compiled := &CompiledController{ PluginID: item.plugin.ID(), Controller: item.ctl, List: list, Form: form, } + if err := BindWritableFields(compiled); err != nil { + return nil, err + } + byID[id] = compiled } return &Registry{byID: byID}, nil }