package lagoon import ( "context" "fmt" "math" "math/big" "reflect" "regexp" "strconv" "strings" "git.golem15.com/golem15/summercms/modules/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 { msgs, err := validateField(ctx, tx, model, field, rule, values, 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, values map[string]any, tr *phrasebook.Translator) ([]string, error) { val := values[field] 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": if !isLaravelBoolean(val) { return []string{validateMessage(ctx, tr, "boolean", field, nil)}, nil } case "email": tags = append(tags, "email") case "confirmed": if fmt.Sprint(val) != fmt.Sprint(values[field+"_confirmation"]) { return []string{validateMessage(ctx, tr, "confirmed", field, nil)}, nil } case "different": if fmt.Sprint(val) == fmt.Sprint(values[arg]) { return []string{validateMessage(ctx, tr, "different", field, map[string]string{"other": arg})}, nil } case "mimes": got := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(fmt.Sprint(val))), ".") match := false for _, ext := range strings.Split(arg, ",") { want := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(ext)), ".") if got != "" && got != "" && got == want { match = true break } } if !match { return []string{validateMessage(ctx, tr, "mimes", field, nil)}, nil } 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 msg, failed := numericRangeMessage(ctx, tr, field, s, rangeMin, rangeMax, minArg == "" && betweenMin != "", maxArg == "" && betweenMax != ""); failed { return []string{msg}, nil } tags = withoutTag(tags, "numeric") } if required && isEmptyValue(val) { return []string{validateMessage(ctx, tr, "required", field, nil)}, nil } usedBetween := betweenMin != "" && betweenMax != "" if len(tags) > 0 { msgs := make([]string, 0, len(tags)) for _, t := range tags { name, _, _ := strings.Cut(t, "=") if name == "required" || name == "min" || name == "max" { continue } if err := validateOnce.Var(val, t); err != nil { ruleName := name if name == "oneof" { ruleName = "oneof" } msgs = append(msgs, validateMessage(ctx, tr, ruleName, field, map[string]string{ "min": betweenMin, "max": betweenMax, })) } } if usedBetween && !numericRange { n := 0 if s, ok := val.(string); ok { n = len(s) } else if !isEmptyValue(val) { n = len(fmt.Sprint(val)) } minN, maxN := atoiOr(betweenMin, 0), atoiOr(betweenMax, 0) if n < minN || (maxN > 0 && n > maxN) { msgs = append(msgs, validateMessage(ctx, tr, "between", field, map[string]string{ "min": betweenMin, "max": betweenMax, })) } } if len(msgs) > 0 { return msgs, 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 } // numericRangeMessage checks a numeric value against its bounds and answers // a failure with the message of the bound that failed: min below the lower // bound, max above the upper, and Laravel's numeric between message when // that bound came from between. func numericRangeMessage(ctx context.Context, tr *phrasebook.Translator, field, val, lo, hi string, loFromBetween, hiFromBetween bool) (string, bool) { r := new(big.Rat) if _, ok := r.SetString(val); !ok { return validateMessage(ctx, tr, "max", field, map[string]string{"max": hi, "min": lo}), true } below, above := false, false if lo != "" { if m, ok := new(big.Rat).SetString(lo); ok && r.Cmp(m) < 0 { below = true } } if hi != "" { if m, ok := new(big.Rat).SetString(hi); ok && r.Cmp(m) > 0 { above = true } } params := map[string]string{"min": lo, "max": hi} switch { case (below && loFromBetween) || (above && hiFromBetween): return validateMessage(ctx, tr, "between.numeric", field, params), true case below: return validateMessage(ctx, tr, "min", field, params), true case above: return validateMessage(ctx, tr, "max", field, params), true } return "", false } 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 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 validateMessage(ctx context.Context, tr *phrasebook.Translator, rule, field string, params map[string]string) string { if params == nil { params = map[string]string{} } params["attribute"] = laravelAttribute(field) key := "lagoon::validate." + rule if rule == "between.numeric" { key = "lagoon::validation.between.numeric" } if tr != nil { s := tr.Get(ctx, key, params) if s != "" && s != key { return s } } attr := params["attribute"] switch rule { case "required": return "The " + attr + " field is required." case "integer": return "The " + attr + " must be an integer." case "numeric": return "The " + attr + " must be a number." case "unique": return "The " + attr + " has already been taken." case "max": return "The " + attr + " may not be greater than " + params["max"] + "." case "min": return "The " + attr + " must be at least " + params["min"] + "." case "oneof": return "The selected " + attr + " is invalid." case "email": return "The " + attr + " must be a valid email address." case "confirmed": return "The " + attr + " confirmation does not match." case "different": return "The " + attr + " and " + params["other"] + " must be different." case "mimes": return "The " + attr + " must be a file of the allowed types." case "between": return "The " + attr + " must be between " + params["min"] + " and " + params["max"] + " characters." case "between.numeric": return "The " + attr + " must be between " + params["min"] + " and " + params["max"] + "." case "boolean": return "The " + attr + " field must be true or false." default: return "The " + attr + " is invalid." } } func laravelAttribute(field string) string { return strings.ReplaceAll(field, "_", " ") } func isLaravelBoolean(val any) bool { switch v := val.(type) { case bool: return true case string: switch strings.ToLower(strings.TrimSpace(v)) { case "0", "1", "true", "false": return true } return false case float64: return v == 0 || v == 1 case int: return v == 0 || v == 1 default: s := strings.TrimSpace(fmt.Sprint(val)) return s == "0" || s == "1" || s == "true" || s == "false" } } func atoiOr(s string, fallback int) int { n, err := strconv.Atoi(strings.TrimSpace(s)) if err != nil { return fallback } return n }