mirror of
https://github.com/go-gitea/gitea.git
synced 2026-10-03 06:10:49 +00:00
refactor: move go-chi/binding into Gitea (#39528)
The `gitea.com/go-chi/binding` package only exists for Gitea, so it moves into `modules/web/binding` to fix its bugs directly. Split out of https://github.com/go-gitea/gitea/pull/39504. - GET and HEAD always bind the query - JSON `null` slice elements and nested `TrimSpace` fields bind correctly - Integer fields reject out-of-range values instead of wrapping - An empty JSON body binds nothing and an unknown binding rule is an error Co-authored-by: wxiaoguang <wxiaoguang@gmail.com> Co-authored-by: bircni <bircni@icloud.com>
This commit is contained in:
@@ -0,0 +1,391 @@
|
||||
// Copyright 2014 Martini Authors
|
||||
// Copyright 2014 The Macaron Authors
|
||||
// Copyright 2020 The Gitea Authors
|
||||
// Copyright 2026 The Gitea Authors. All rights reserved.
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package binding binds form, multipart and JSON request data to structs and validates them by their "binding" tags.
|
||||
package binding
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"unicode/utf8"
|
||||
|
||||
"gitea.dev/modules/json"
|
||||
"gitea.dev/modules/util"
|
||||
)
|
||||
|
||||
const (
|
||||
errContentType = "ContentTypeError"
|
||||
errDeserialization = "DeserializationError"
|
||||
errTypeCast = "TypeCastError"
|
||||
errRule = "RuleError"
|
||||
|
||||
ErrRequired = "RequiredError"
|
||||
ErrAlphaDashDot = "AlphaDashDotError"
|
||||
ErrMinSize = "MinSizeError"
|
||||
ErrMaxSize = "MaxSizeError"
|
||||
ErrRange = "RangeError"
|
||||
ErrIn = "InError"
|
||||
ErrInclude = "IncludeError"
|
||||
)
|
||||
|
||||
const multipartMaxMemory = 10 * 1024 * 1024
|
||||
|
||||
type (
|
||||
Errors []Error
|
||||
|
||||
Error struct {
|
||||
FieldNames []string
|
||||
Classification string
|
||||
Message string
|
||||
}
|
||||
)
|
||||
|
||||
func (e Error) Error() string {
|
||||
return e.Message
|
||||
}
|
||||
|
||||
func (e *Errors) addDeserializationError(err error) {
|
||||
*e = append(*e, Error{Classification: errDeserialization, Message: err.Error()})
|
||||
}
|
||||
|
||||
func (e *Errors) addOptional(err *Error) {
|
||||
if err != nil {
|
||||
*e = append(*e, *err)
|
||||
}
|
||||
}
|
||||
|
||||
func newFieldError(field reflect.StructField, classification, message string) *Error {
|
||||
return &Error{FieldNames: []string{field.Name}, Classification: classification, Message: message}
|
||||
}
|
||||
|
||||
type ValidationField struct {
|
||||
StructField reflect.StructField
|
||||
reflectValue reflect.Value
|
||||
ruleArgs []string
|
||||
}
|
||||
|
||||
func (f *ValidationField) valueAsString() string {
|
||||
return fmt.Sprint(reflect.Indirect(f.reflectValue).Interface())
|
||||
}
|
||||
|
||||
func (f *ValidationField) ValueMustString() string {
|
||||
value := reflect.Indirect(f.reflectValue)
|
||||
if value.Kind() != reflect.String {
|
||||
panic("field value must be a string")
|
||||
}
|
||||
return value.String()
|
||||
}
|
||||
|
||||
func (f *ValidationField) valueSize() int {
|
||||
value := reflect.Indirect(f.reflectValue)
|
||||
switch value.Kind() {
|
||||
case reflect.String:
|
||||
return utf8.RuneCountInString(value.String())
|
||||
case reflect.Slice:
|
||||
return value.Len()
|
||||
}
|
||||
panic("unsupported type: " + value.Kind().String())
|
||||
}
|
||||
|
||||
func (f *ValidationField) assignValue(newValue any) {
|
||||
if f.reflectValue.Kind() == reflect.Pointer {
|
||||
ptr := reflect.New(f.StructField.Type.Elem())
|
||||
ptr.Elem().Set(reflect.ValueOf(newValue).Convert(ptr.Elem().Type()))
|
||||
f.reflectValue.Set(ptr)
|
||||
return
|
||||
}
|
||||
f.reflectValue.Set(reflect.ValueOf(newValue).Convert(f.reflectValue.Type()))
|
||||
}
|
||||
|
||||
type RuleValidator func(ctx context.Context, field *ValidationField) *Error
|
||||
|
||||
type ruleValidatorItem struct {
|
||||
forZeroValue bool
|
||||
validatorFn RuleValidator
|
||||
}
|
||||
|
||||
type Binder struct {
|
||||
rules map[string]ruleValidatorItem
|
||||
}
|
||||
|
||||
func NewBinder() *Binder {
|
||||
binder := &Binder{rules: map[string]ruleValidatorItem{}}
|
||||
binder.AddRuleNonZero("TrimSpace", func(_ context.Context, field *ValidationField) *Error {
|
||||
stringType := reflect.TypeFor[string]()
|
||||
value := reflect.Indirect(field.reflectValue)
|
||||
if !value.CanConvert(stringType) {
|
||||
return newFieldError(field.StructField, errTypeCast, "TrimSpace")
|
||||
}
|
||||
field.assignValue(strings.TrimSpace(value.Convert(stringType).String()))
|
||||
return nil
|
||||
})
|
||||
binder.rules["Required"] = ruleValidatorItem{forZeroValue: true, validatorFn: func(_ context.Context, field *ValidationField) *Error {
|
||||
// a pointer field is optional, so "Required" only applies once a value is provided
|
||||
if field.reflectValue.Kind() == reflect.Pointer && field.reflectValue.IsNil() {
|
||||
return nil
|
||||
}
|
||||
return newFieldError(field.StructField, ErrRequired, "Required")
|
||||
}}
|
||||
binder.AddRuleNonZero("AlphaDashDot", func(_ context.Context, field *ValidationField) *Error {
|
||||
if nonAlphaDashDotPattern().MatchString(field.ValueMustString()) {
|
||||
return newFieldError(field.StructField, ErrAlphaDashDot, "AlphaDashDot")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
binder.AddRuleNonZero("MinSize", func(_ context.Context, field *ValidationField) *Error {
|
||||
minSize, _ := strconv.Atoi(field.ruleArgs[0])
|
||||
if field.valueSize() < minSize {
|
||||
return newFieldError(field.StructField, ErrMinSize, "MinSize")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
binder.AddRuleNonZero("MaxSize", func(_ context.Context, field *ValidationField) *Error {
|
||||
maxSize, _ := strconv.Atoi(field.ruleArgs[0])
|
||||
if field.valueSize() > maxSize {
|
||||
return newFieldError(field.StructField, ErrMaxSize, "MaxSize")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
binder.AddRuleNonZero("Range", func(_ context.Context, field *ValidationField) *Error {
|
||||
value, _ := strconv.Atoi(field.valueAsString())
|
||||
minValue, _ := strconv.Atoi(field.ruleArgs[0])
|
||||
maxValue, _ := strconv.Atoi(field.ruleArgs[1])
|
||||
if value < minValue || value > maxValue {
|
||||
return newFieldError(field.StructField, ErrRange, "Range")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
binder.AddRuleNonZero("In", func(_ context.Context, field *ValidationField) *Error {
|
||||
if !slices.Contains(field.ruleArgs, field.valueAsString()) {
|
||||
return newFieldError(field.StructField, ErrIn, "In")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
binder.AddRuleNonZero("Include", func(_ context.Context, field *ValidationField) *Error {
|
||||
if !strings.Contains(field.ValueMustString(), strings.Join(field.ruleArgs, ",")) {
|
||||
return newFieldError(field.StructField, ErrInclude, "Include")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
return binder
|
||||
}
|
||||
|
||||
func (b *Binder) AddRuleNonZero(name string, ruleValidator RuleValidator) {
|
||||
b.rules[name] = ruleValidatorItem{validatorFn: ruleValidator}
|
||||
}
|
||||
|
||||
func (b *Binder) Bind(req *http.Request, obj any) Errors {
|
||||
ensurePointer(obj)
|
||||
contentType := req.Header.Get("Content-Type")
|
||||
if req.Method == http.MethodGet || req.Method == http.MethodHead ||
|
||||
(contentType == "" && req.Method != http.MethodPost && req.Method != http.MethodPut) {
|
||||
return b.bindForm(req, obj)
|
||||
}
|
||||
switch {
|
||||
case strings.Contains(contentType, "form-urlencoded"):
|
||||
return b.bindForm(req, obj)
|
||||
case strings.Contains(contentType, "multipart/form-data"):
|
||||
return b.bindMultipartForm(req, obj)
|
||||
case strings.Contains(contentType, "json"):
|
||||
return b.bindJSON(req, obj)
|
||||
}
|
||||
return Errors{{Classification: errContentType, Message: "Unsupported Content-Type"}}
|
||||
}
|
||||
|
||||
func (b *Binder) bindForm(req *http.Request, formStruct any) (errs Errors) {
|
||||
if err := req.ParseForm(); err != nil {
|
||||
errs.addDeserializationError(err)
|
||||
}
|
||||
errs = mapForm(reflect.ValueOf(formStruct), req.Form, nil, errs)
|
||||
return append(errs, b.Validate(req.Context(), formStruct)...)
|
||||
}
|
||||
|
||||
func (b *Binder) bindMultipartForm(req *http.Request, formStruct any) (errs Errors) {
|
||||
if err := req.ParseMultipartForm(multipartMaxMemory); err != nil {
|
||||
errs.addDeserializationError(err)
|
||||
}
|
||||
if req.MultipartForm != nil {
|
||||
errs = mapForm(reflect.ValueOf(formStruct), req.MultipartForm.Value, req.MultipartForm.File, errs)
|
||||
}
|
||||
return append(errs, b.Validate(req.Context(), formStruct)...)
|
||||
}
|
||||
|
||||
func (b *Binder) bindJSON(req *http.Request, jsonStruct any) (errs Errors) {
|
||||
err := json.NewDecoder(req.Body).Decode(jsonStruct)
|
||||
if err != nil && !errors.Is(err, io.EOF) { // an empty body binds nothing
|
||||
errs.addDeserializationError(err)
|
||||
}
|
||||
return append(errs, b.Validate(req.Context(), jsonStruct)...)
|
||||
}
|
||||
|
||||
func (b *Binder) Validate(ctx context.Context, obj any) Errors {
|
||||
ensurePointer(obj)
|
||||
return b.validateStruct(ctx, nil, reflect.ValueOf(obj).Elem())
|
||||
}
|
||||
|
||||
func (b *Binder) validateStruct(ctx context.Context, errs Errors, structValue reflect.Value) Errors {
|
||||
structType := structValue.Type()
|
||||
for i := range structType.NumField() {
|
||||
field := structType.Field(i)
|
||||
fieldValue := structValue.Field(i)
|
||||
if field.Tag.Get("form") == "-" || !fieldValue.CanInterface() {
|
||||
continue
|
||||
}
|
||||
if field.Type.Kind() == reflect.Struct ||
|
||||
(field.Type.Kind() == reflect.Pointer && !fieldValue.IsNil() && field.Type.Elem().Kind() == reflect.Struct) {
|
||||
errs = b.validateStruct(ctx, errs, reflect.Indirect(fieldValue))
|
||||
continue
|
||||
}
|
||||
errs = b.validateField(ctx, errs, &ValidationField{StructField: field, reflectValue: fieldValue})
|
||||
}
|
||||
return errs
|
||||
}
|
||||
|
||||
func (b *Binder) validateField(ctx context.Context, errs Errors, field *ValidationField) Errors {
|
||||
if field.reflectValue.Kind() == reflect.Slice {
|
||||
for i := range field.reflectValue.Len() {
|
||||
if elem := reflect.Indirect(field.reflectValue.Index(i)); elem.Kind() == reflect.Struct {
|
||||
errs = b.validateStruct(ctx, errs, elem)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for rule := range strings.SplitSeq(field.StructField.Tag.Get("binding"), ";") {
|
||||
rule = strings.TrimSpace(rule)
|
||||
if rule == "" {
|
||||
continue
|
||||
}
|
||||
ruleName, ruleArgs, _ := strings.Cut(rule, "(")
|
||||
field.ruleArgs = nil
|
||||
if ruleArgs != "" {
|
||||
field.ruleArgs = strings.Split(strings.TrimSuffix(ruleArgs, ")"), ",")
|
||||
}
|
||||
item, ok := b.rules[ruleName]
|
||||
if !ok {
|
||||
panic(fmt.Sprintf("Invalid binding rule: %q", ruleName))
|
||||
}
|
||||
value := reflect.Indirect(field.reflectValue)
|
||||
if item.forZeroValue != (!value.IsValid() || value.IsZero()) {
|
||||
continue
|
||||
}
|
||||
if err := item.validatorFn(ctx, field); err != nil {
|
||||
return append(errs, *err)
|
||||
}
|
||||
}
|
||||
return errs
|
||||
}
|
||||
|
||||
var nonAlphaDashDotPattern = sync.OnceValue(func() *regexp.Regexp {
|
||||
return regexp.MustCompile(`[^\w-.]`)
|
||||
})
|
||||
|
||||
func mapForm(formStruct reflect.Value, form map[string][]string, formFiles map[string][]*multipart.FileHeader, errs Errors) Errors {
|
||||
formStruct = reflect.Indirect(formStruct)
|
||||
structType := formStruct.Type()
|
||||
|
||||
for fieldIdx := range structType.NumField() {
|
||||
typeField := structType.Field(fieldIdx)
|
||||
fieldValue := formStruct.Field(fieldIdx)
|
||||
|
||||
if typeField.Type.Kind() == reflect.Pointer && typeField.Anonymous {
|
||||
fieldValue.Set(reflect.New(typeField.Type.Elem()))
|
||||
errs = mapForm(fieldValue.Elem(), form, formFiles, errs)
|
||||
if fieldValue.Elem().IsZero() {
|
||||
fieldValue.SetZero()
|
||||
}
|
||||
} else if typeField.Type.Kind() == reflect.Struct {
|
||||
errs = mapForm(fieldValue, form, formFiles, errs)
|
||||
}
|
||||
|
||||
inputFieldName := typeField.Tag.Get("form")
|
||||
if inputFieldName == "-" || !typeField.IsExported() {
|
||||
continue
|
||||
}
|
||||
if inputFieldName == "" {
|
||||
inputFieldName = util.ToSnakeCase(typeField.Name)
|
||||
}
|
||||
|
||||
if inputValues, exists := form[inputFieldName]; exists {
|
||||
if fieldValue.Kind() == reflect.Slice && len(inputValues) > 0 {
|
||||
slice := reflect.MakeSlice(fieldValue.Type(), len(inputValues), len(inputValues))
|
||||
for elemIdx, inputValue := range inputValues {
|
||||
errs.addOptional(setWithProperType(typeField, inputValue, slice.Index(elemIdx)))
|
||||
}
|
||||
fieldValue.Set(slice)
|
||||
} else {
|
||||
errs.addOptional(setWithProperType(typeField, inputValues[0], fieldValue))
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
inputFiles, exists := formFiles[inputFieldName]
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
fileHeaderType := reflect.TypeFor[*multipart.FileHeader]()
|
||||
if fieldValue.Kind() == reflect.Slice && len(inputFiles) > 0 && fieldValue.Type().Elem() == fileHeaderType {
|
||||
fieldValue.Set(reflect.ValueOf(slices.Clone(inputFiles)))
|
||||
} else if fieldValue.Type() == fileHeaderType {
|
||||
fieldValue.Set(reflect.ValueOf(inputFiles[0]))
|
||||
}
|
||||
}
|
||||
return errs
|
||||
}
|
||||
|
||||
func setWithProperType(structField reflect.StructField, val string, fieldValue reflect.Value) *Error {
|
||||
switch fieldValue.Kind() {
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
intVal, err := strconv.ParseInt(cmp.Or(val, "0"), 10, fieldValue.Type().Bits())
|
||||
if err != nil {
|
||||
return newFieldError(structField, errTypeCast, "Value could not be parsed as integer")
|
||||
}
|
||||
fieldValue.SetInt(intVal)
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
||||
uintVal, err := strconv.ParseUint(cmp.Or(val, "0"), 10, fieldValue.Type().Bits())
|
||||
if err != nil {
|
||||
return newFieldError(structField, errTypeCast, "Value could not be parsed as unsigned integer")
|
||||
}
|
||||
fieldValue.SetUint(uintVal)
|
||||
case reflect.Bool:
|
||||
if val == "on" {
|
||||
fieldValue.SetBool(true)
|
||||
break
|
||||
}
|
||||
boolVal, err := strconv.ParseBool(cmp.Or(val, "false"))
|
||||
if err != nil {
|
||||
return newFieldError(structField, errTypeCast, "Value could not be parsed as boolean")
|
||||
}
|
||||
fieldValue.SetBool(boolVal)
|
||||
case reflect.String:
|
||||
fieldValue.SetString(val)
|
||||
case reflect.Pointer:
|
||||
newValue := reflect.New(fieldValue.Type().Elem())
|
||||
if err := setWithProperType(structField, val, newValue.Elem()); err != nil {
|
||||
return err
|
||||
}
|
||||
fieldValue.Set(newValue)
|
||||
default:
|
||||
return newFieldError(structField, errDeserialization, "unsupported type: "+fieldValue.Kind().String())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensurePointer(obj any) {
|
||||
if reflect.TypeOf(obj).Kind() != reflect.Pointer {
|
||||
panic("Pointers are only accepted as binding models")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user