This commit is contained in:
oneao committed 2025-11-05 22:25:09 +08:00
1 parent c7a1f83c67
commit 6849579652
19 files changed
+600 -249

No files matched your search

@@ -1,42 +1,65 @@
package validate
import (
"base-go-v2/internal/errs"
"errors"
"reflect"
"fmt"
"strings"
)
// ValidateNotEmpty 校验结构体中指定字段不能为空
func ValidateNotEmpty(s interface{}, fields ...string) error {
v := reflect.ValueOf(s)
if v.Kind() == reflect.Ptr {
v = v.Elem()
// ValidationError 用于收集多个字段错误,同时实现 error 接口
type ValidationError struct {
Errors map[string]string
}
// Error 实现 error 接口
func (v *ValidationError) Error() string {
msgs := make([]string, 0, len(v.Errors))
for field, err := range v.Errors {
msgs = append(msgs, fmt.Sprintf("%s: %s", field, err))
}
return strings.Join(msgs, "; ")
}
// NotEmpty 校验 map[string]interface{} 中指定字段不能为空
func NotEmpty(data map[string]interface{}, fields ...string) error {
errs := make(map[string]string)
for _, field := range fields {
f := v.FieldByName(field)
if !f.IsValid() {
return errors.New("field " + field + " does not exist")
}
switch f.Kind() {
case reflect.String:
if strings.TrimSpace(f.String()) == "" {
return errs.NewValidationError(field, "不能为空")
}
case reflect.Slice, reflect.Array, reflect.Map:
if f.Len() == 0 {
return errs.NewValidationError(field, "不能为空")
}
case reflect.Ptr, reflect.Interface:
if f.IsNil() {
return errs.NewValidationError(field, "不能为空")
}
default:
continue
}
collectErrors(data[field], field, errs)
}
if len(errs) > 0 {
return &ValidationError{Errors: errs}
}
return nil
}
// collectErrors 递归收集字段错误
func collectErrors(v interface{}, path string, errs map[string]string) {
if v == nil {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
return
}
switch val := v.(type) {
case string:
if strings.TrimSpace(val) == "" {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
}
case []interface{}:
if len(val) == 0 {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
}
for i, elem := range val {
collectErrors(elem, fmt.Sprintf("%s[%d]", path, i), errs)
}
case map[string]interface{}:
if len(val) == 0 {
errs[path] = fmt.Sprintf("字段 %s 不能为空", path)
}
for k, elem := range val {
collectErrors(elem, fmt.Sprintf("%s.%s", path, k), errs)
}
default:
// 其他类型暂不校验
}
}