package lagoon import ( "context" "fmt" "math" "math/big" "reflect" "regexp" "strconv" "strings" "git.golem15.com/golem15/summercms/phrasebook" "github.com/go-playground/validator/v10" "gorm.io/gorm" ) var identName = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_]*$`) var validateOnce = validator.New(validator.WithRequiredStructEnabled()) // Validate translates Laravel-style rule strings onto validator.Var() plus a // unique:table DB check. Unrecognized tokens fail loudly. The returned map is // the Laravel-shaped errors object (field -> messages); the HTTP envelope is // a later-phase concern. func Validate(ctx context.Context, tx *gorm.DB, model any, rules map[string]string, values map[string]any, tr *phrasebook.Translator) (map[string][]string, error) { out := map[string][]string{} for field, rule := range rules { val, _ := values[field] msgs, err := validateField(ctx, tx, model, field, rule, val, tr) if err != nil { return nil, err } if len(msgs) > 0 { out[field] = msgs } } if len(out) == 0 { return nil, nil } return out, nil } func validateField(ctx context.Context, tx *gorm.DB, model any, field, rule string, val any, tr *phrasebook.Translator) ([]string, error) { tokens := splitRule(rule) nullable := false required := false var tags []string var uniqueTable string var betweenMin, betweenMax, minArg, maxArg string hasNumeric := false hasInteger := false for _, tok := range tokens { name, arg, _ := strings.Cut(tok, ":") switch name { case "nullable": nullable = true case "required": required = true tags = append(tags, "required") case "integer": hasInteger = true if !isIntegerValue(val) && !isEmptyValue(val) { return []string{validateMessage(ctx, tr, "integer", field, nil)}, nil } case "numeric": hasNumeric = true tags = append(tags, "numeric") case "between": x, y, ok := strings.Cut(arg, ",") if !ok { return nil, fmt.Errorf("lagoon: unrecognized validation rule %q", tok) } betweenMin, betweenMax = x, y case "min": minArg = arg case "max": maxArg = arg case "in": tags = append(tags, oneofTag(strings.Split(arg, ","))) case "unique": uniqueTable = arg case "boolean": // Go's bool field type already enforces this; treat as a type-check no-op. default: return nil, fmt.Errorf("lagoon: unrecognized validation rule %q", tok) } } if nullable && isEmptyValue(val) && !required { return nil, nil } numericRange := hasNumeric || hasInteger var rangeMin, rangeMax string if numericRange { if betweenMin != "" { rangeMin, rangeMax = betweenMin, betweenMax } if minArg != "" { rangeMin = minArg } if maxArg != "" { rangeMax = maxArg } } else { if betweenMin != "" { tags = append(tags, "min="+betweenMin, "max="+betweenMax) } if minArg != "" { tags = append(tags, "min="+minArg) } if maxArg != "" { tags = append(tags, "max="+maxArg) } } if numericRange && (rangeMin != "" || rangeMax != "") { s, ok := numericString(val) if !ok { ruleName := "numeric" if hasInteger { ruleName = "integer" } return []string{validateMessage(ctx, tr, ruleName, field, nil)}, nil } if !moneyInRange(s, rangeMin, rangeMax) { return []string{validateMessage(ctx, tr, "max", field, map[string]string{"max": rangeMax, "min": rangeMin})}, nil } tags = withoutTag(tags, "numeric") } if len(tags) > 0 { tag := strings.Join(tags, ",") if err := validateOnce.Var(val, tag); err != nil { ruleName := "required" if !required { ruleName = firstNonOmit(tags) } return []string{validateMessage(ctx, tr, ruleName, field, nil)}, nil } } if uniqueTable != "" { ok, err := uniqueOK(tx, model, uniqueTable, field, val) if err != nil { return nil, err } if !ok { return []string{validateMessage(ctx, tr, "unique", field, nil)}, nil } } return nil, nil } func splitRule(rule string) []string { var out []string for _, p := range strings.Split(rule, "|") { p = strings.TrimSpace(p) if p != "" { out = append(out, p) } } return out } func oneofTag(vals []string) string { parts := make([]string, 0, len(vals)) for _, v := range vals { v = strings.TrimSpace(v) if v == "" { continue } if strings.ContainsAny(v, " \t\"'") { parts = append(parts, "'"+v+"'") } else { parts = append(parts, v) } } return "oneof=" + strings.Join(parts, " ") } func isEmptyValue(val any) bool { if val == nil { return true } rv := reflect.ValueOf(val) switch rv.Kind() { case reflect.Ptr, reflect.Interface: return rv.IsNil() case reflect.String, reflect.Slice, reflect.Map: return rv.Len() == 0 default: return false } } // isIntegerValue reports whether val is an integer in Laravel's sense. It // dereferences pointers (model fields such as *int) and accepts whole // floating-point numbers, the shape encoding/json gives a map[string]any. func isIntegerValue(val any) bool { if val == nil { return true } rv := reflect.ValueOf(val) for rv.Kind() == reflect.Ptr || rv.Kind() == reflect.Interface { if rv.IsNil() { return true } rv = rv.Elem() } switch rv.Kind() { case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: return true case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: return true case reflect.Float32, reflect.Float64: f := rv.Float() return f == math.Trunc(f) && !math.IsInf(f, 0) case reflect.String: s := rv.String() if s == "" { return true } _, ok := new(big.Int).SetString(s, 10) return ok default: return false } } func numericString(val any) (string, bool) { if val == nil { return "", false } rv := reflect.ValueOf(val) for rv.Kind() == reflect.Ptr || rv.Kind() == reflect.Interface { if rv.IsNil() { return "", false } rv = rv.Elem() } switch rv.Kind() { case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: return strconv.FormatInt(rv.Int(), 10), true case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: return strconv.FormatUint(rv.Uint(), 10), true case reflect.Float32, reflect.Float64: return strconv.FormatFloat(rv.Float(), 'f', -1, 64), true case reflect.String: s := strings.TrimSpace(rv.String()) if s == "" { return "", false } return s, true default: return "", false } } func moneyInRange(val any, min, max string) bool { s, ok := numericString(val) if !ok { s = strings.TrimSpace(fmt.Sprint(val)) if s == "" || s == "" { return true } } r := new(big.Rat) if _, ok := r.SetString(s); !ok { return false } if min != "" { m := new(big.Rat) if _, ok := m.SetString(min); ok && r.Cmp(m) < 0 { return false } } if max != "" { m := new(big.Rat) if _, ok := m.SetString(max); ok && r.Cmp(m) > 0 { return false } } return true } func uniqueOK(tx *gorm.DB, model any, table, column string, val any) (bool, error) { if tx == nil { return false, fmt.Errorf("lagoon: unique:%s requires a database handle", table) } if !identName.MatchString(table) || !identName.MatchString(column) { return false, fmt.Errorf("lagoon: unique identifier %q.%q is not safe", table, column) } if isEmptyValue(val) { return true, nil } q := tx.Table(table).Where(column+" = ?", val) if tx.Migrator().HasColumn(table, "deleted_at") { q = q.Where("deleted_at IS NULL") } if id := modelUintID(model); id != 0 && tx.Migrator().HasColumn(table, "id") { q = q.Where("id <> ?", id) } var n int64 if err := q.Count(&n).Error; err != nil { return false, err } return n == 0, nil } func modelUintID(model any) uint { if model == nil { return 0 } rv := reflect.ValueOf(model) if rv.Kind() == reflect.Ptr { if rv.IsNil() { return 0 } rv = rv.Elem() } if rv.Kind() != reflect.Struct { return 0 } f := rv.FieldByName("ID") if !f.IsValid() { return 0 } switch f.Kind() { case reflect.Uint, reflect.Uint32, reflect.Uint64: return uint(f.Uint()) case reflect.Int, reflect.Int32, reflect.Int64: if f.Int() < 0 { return 0 } return uint(f.Int()) } return 0 } func withoutTag(tags []string, name string) []string { out := tags[:0] for _, t := range tags { n, _, _ := strings.Cut(t, "=") if n != name { out = append(out, t) } } return out } func firstNonOmit(tags []string) string { for _, t := range tags { name, _, _ := strings.Cut(t, "=") if name != "omitempty" && name != "required" { return name } } return "invalid" } func validateMessage(ctx context.Context, tr *phrasebook.Translator, rule, field string, params map[string]string) string { if params == nil { params = map[string]string{} } params["attribute"] = field key := "lagoon.validate." + rule if tr != nil { s := tr.Get(ctx, key, params) if s != "" && s != key { return s } } switch rule { case "required": return "The " + field + " field is required." case "integer": return "The " + field + " must be an integer." case "numeric": return "The " + field + " must be a number." case "unique": return "The " + field + " has already been taken." case "max": return "The " + field + " may not be greater than " + params["max"] + "." case "min": return "The " + field + " must be at least " + params["min"] + "." case "oneof": return "The selected " + field + " is invalid." default: return "The " + field + " is invalid." } }