init
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
package validation
|
||||
|
||||
import (
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// 默认加密成本(bcrypt 推荐值)
|
||||
const defaultCost = 12
|
||||
|
||||
// HashPassword 对密码进行加密
|
||||
func HashPassword(password string) (string, error) {
|
||||
return HashPasswordWithCost(password, defaultCost)
|
||||
}
|
||||
|
||||
// HashPasswordWithCost 使用指定成本对密码进行加密
|
||||
// cost 范围: 4-31,值越大越安全但越慢
|
||||
func HashPasswordWithCost(password string, cost int) (string, error) {
|
||||
if cost < bcrypt.MinCost {
|
||||
cost = bcrypt.MinCost
|
||||
}
|
||||
if cost > bcrypt.MaxCost {
|
||||
cost = bcrypt.MaxCost
|
||||
}
|
||||
|
||||
bytes, err := bcrypt.GenerateFromPassword([]byte(password), cost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(bytes), nil
|
||||
}
|
||||
|
||||
// CheckPassword 验证密码是否匹配
|
||||
// hashedPassword 是加密后的密码,password 是明文密码
|
||||
func CheckPassword(hashedPassword, password string) bool {
|
||||
err := bcrypt.CompareHashAndPassword([]byte(hashedPassword), []byte(password))
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// CheckPasswordAndUpgrade 验证密码并在需要时升级加密成本
|
||||
// 返回:是否匹配、是否需要升级、升级后的密码、错误
|
||||
func CheckPasswordAndUpgrade(hashedPassword, password string, targetCost int) (match bool, needUpgrade bool, newHash string, err error) {
|
||||
// 验证密码
|
||||
err = bcrypt.CompareHashAndPassword([]byte(hashedPassword), []byte(password))
|
||||
if err != nil {
|
||||
return false, false, "", err
|
||||
}
|
||||
|
||||
// 检查是否需要升级
|
||||
cost, err := bcrypt.Cost([]byte(hashedPassword))
|
||||
if err != nil {
|
||||
return true, false, "", nil
|
||||
}
|
||||
|
||||
if cost < targetCost {
|
||||
newHash, err = HashPasswordWithCost(password, targetCost)
|
||||
if err != nil {
|
||||
return true, false, "", err
|
||||
}
|
||||
return true, true, newHash, nil
|
||||
}
|
||||
|
||||
return true, false, "", nil
|
||||
}
|
||||
|
||||
// GetPasswordCost 获取加密密码的成本
|
||||
func GetPasswordCost(hashedPassword string) (int, error) {
|
||||
return bcrypt.Cost([]byte(hashedPassword))
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
package validation
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"unicode"
|
||||
)
|
||||
|
||||
// PasswordConfig 密码验证配置
|
||||
type PasswordConfig struct {
|
||||
MinLength int // 最小长度
|
||||
MaxLength int // 最大长度
|
||||
RequireUpper bool // 需要大写字母
|
||||
RequireLower bool // 需要小写字母
|
||||
RequireDigit bool // 需要数字
|
||||
RequireSpecial bool // 需要特殊字符
|
||||
}
|
||||
|
||||
// DefaultPasswordConfig 默认密码配置
|
||||
var DefaultPasswordConfig = PasswordConfig{
|
||||
MinLength: 8, // 最少8位
|
||||
MaxLength: 128, // 最多128位
|
||||
RequireUpper: true,
|
||||
RequireLower: true,
|
||||
RequireDigit: true,
|
||||
RequireSpecial: false,
|
||||
}
|
||||
|
||||
// ValidatePassword 验证密码强度
|
||||
// 返回:是否有效,错误信息
|
||||
func ValidatePassword(password string) (bool, string) {
|
||||
return ValidatePasswordWithConfig(password, DefaultPasswordConfig)
|
||||
}
|
||||
|
||||
// ValidatePasswordWithConfig 使用指定配置验证密码强度
|
||||
func ValidatePasswordWithConfig(password string, config PasswordConfig) (bool, string) {
|
||||
length := len(password)
|
||||
|
||||
// 检查最小长度
|
||||
if length < config.MinLength {
|
||||
return false, "密码长度不能少于" + strconv.Itoa(config.MinLength) + "位"
|
||||
}
|
||||
|
||||
// 检查最大长度
|
||||
if length > config.MaxLength {
|
||||
return false, "密码长度不能超过" + strconv.Itoa(config.MaxLength) + "位"
|
||||
}
|
||||
|
||||
var hasUpper, hasLower, hasDigit, hasSpecial bool
|
||||
|
||||
for _, char := range password {
|
||||
switch {
|
||||
case unicode.IsUpper(char):
|
||||
hasUpper = true
|
||||
case unicode.IsLower(char):
|
||||
hasLower = true
|
||||
case unicode.IsDigit(char):
|
||||
hasDigit = true
|
||||
case unicode.IsPunct(char) || unicode.IsSymbol(char):
|
||||
hasSpecial = true
|
||||
}
|
||||
}
|
||||
|
||||
// 检查大写字母
|
||||
if config.RequireUpper && !hasUpper {
|
||||
return false, "密码必须包含大写字母"
|
||||
}
|
||||
|
||||
// 检查小写字母
|
||||
if config.RequireLower && !hasLower {
|
||||
return false, "密码必须包含小写字母"
|
||||
}
|
||||
|
||||
// 检查数字
|
||||
if config.RequireDigit && !hasDigit {
|
||||
return false, "密码必须包含数字"
|
||||
}
|
||||
|
||||
// 检查特殊字符
|
||||
if config.RequireSpecial && !hasSpecial {
|
||||
return false, "密码必须包含特殊字符"
|
||||
}
|
||||
|
||||
return true, ""
|
||||
}
|
||||
@@ -0,0 +1,403 @@
|
||||
package validation_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/EthanCodeCraft/xlgo-core/validation"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
// ===== Password Tests =====
|
||||
|
||||
func TestValidatePassword(t *testing.T) {
|
||||
tests := []struct {
|
||||
password string
|
||||
valid bool
|
||||
msg string
|
||||
}{
|
||||
{"Abc12345", true, ""}, // 有效
|
||||
{"abc12345", false, "密码必须包含大写字母"}, // 缺大写
|
||||
{"ABC12345", false, "密码必须包含小写字母"}, // 缺小写
|
||||
{"Abcdefgh", false, "密码必须包含数字"}, // 缺数字
|
||||
{"Abc123", false, "密码长度不能少于8位"}, // 太短
|
||||
{"", false, "密码长度不能少于8位"}, // 空
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
valid, msg := validation.ValidatePassword(tt.password)
|
||||
if valid != tt.valid {
|
||||
t.Errorf("ValidatePassword(%s) valid = %v, want %v", tt.password, valid, tt.valid)
|
||||
}
|
||||
if !valid && msg != tt.msg {
|
||||
t.Errorf("ValidatePassword(%s) msg = %s, want %s", tt.password, msg, tt.msg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidatePasswordWithConfig(t *testing.T) {
|
||||
// 自定义配置:不要求特殊字符
|
||||
config := validation.PasswordConfig{
|
||||
MinLength: 6,
|
||||
MaxLength: 20,
|
||||
RequireUpper: false,
|
||||
RequireLower: true,
|
||||
RequireDigit: true,
|
||||
RequireSpecial: false,
|
||||
}
|
||||
|
||||
valid, msg := validation.ValidatePasswordWithConfig("abc123", config)
|
||||
if !valid {
|
||||
t.Errorf("ValidatePasswordWithConfig should be valid: %s", msg)
|
||||
}
|
||||
|
||||
// 要求特殊字符
|
||||
config2 := validation.PasswordConfig{
|
||||
MinLength: 8,
|
||||
MaxLength: 20,
|
||||
RequireUpper: true,
|
||||
RequireLower: true,
|
||||
RequireDigit: true,
|
||||
RequireSpecial: true,
|
||||
}
|
||||
|
||||
valid2, msg2 := validation.ValidatePasswordWithConfig("Abc12345", config2)
|
||||
if valid2 {
|
||||
t.Error("Should require special character")
|
||||
}
|
||||
if msg2 != "密码必须包含特殊字符" {
|
||||
t.Errorf("msg = %s, want '密码必须包含特殊字符'", msg2)
|
||||
}
|
||||
|
||||
// 包含特殊字符
|
||||
valid3, msg3 := validation.ValidatePasswordWithConfig("Abc123!@#", config2)
|
||||
if !valid3 {
|
||||
t.Errorf("Should be valid: %s", msg3)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultPasswordConfig(t *testing.T) {
|
||||
cfg := validation.DefaultPasswordConfig
|
||||
if cfg.MinLength != 8 {
|
||||
t.Errorf("MinLength = %d, want 8", cfg.MinLength)
|
||||
}
|
||||
if cfg.MaxLength != 128 {
|
||||
t.Errorf("MaxLength = %d, want 128", cfg.MaxLength)
|
||||
}
|
||||
if !cfg.RequireUpper || !cfg.RequireLower || !cfg.RequireDigit {
|
||||
t.Error("Default config should require upper, lower, digit")
|
||||
}
|
||||
}
|
||||
|
||||
// ===== Hash Password Tests =====
|
||||
|
||||
func TestHashPassword(t *testing.T) {
|
||||
password := "testPassword123"
|
||||
|
||||
hash, err := validation.HashPassword(password)
|
||||
if err != nil {
|
||||
t.Fatalf("HashPassword error: %v", err)
|
||||
}
|
||||
|
||||
// 验证 hash 不等于原密码
|
||||
if hash == password {
|
||||
t.Error("Hash should not equal original password")
|
||||
}
|
||||
|
||||
// 验证 hash 非空
|
||||
if hash == "" {
|
||||
t.Error("Hash should not be empty")
|
||||
}
|
||||
|
||||
// 验证可以匹配
|
||||
if !validation.CheckPassword(hash, password) {
|
||||
t.Error("CheckPassword should return true for correct password")
|
||||
}
|
||||
|
||||
// 验证错误密码不匹配
|
||||
if validation.CheckPassword(hash, "wrongPassword") {
|
||||
t.Error("CheckPassword should return false for wrong password")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHashPasswordWithCost(t *testing.T) {
|
||||
password := "testPassword123"
|
||||
|
||||
// 测试不同 cost
|
||||
hash, err := validation.HashPasswordWithCost(password, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("HashPasswordWithCost error: %v", err)
|
||||
}
|
||||
|
||||
cost, err := validation.GetPasswordCost(hash)
|
||||
if err != nil {
|
||||
t.Fatalf("GetPasswordCost error: %v", err)
|
||||
}
|
||||
if cost != 10 {
|
||||
t.Errorf("Cost = %d, want 10", cost)
|
||||
}
|
||||
|
||||
// 测试 cost 边界
|
||||
hash2, err := validation.HashPasswordWithCost(password, 3) // 低于 MinCost
|
||||
if err != nil {
|
||||
t.Fatalf("HashPasswordWithCost with low cost error: %v", err)
|
||||
}
|
||||
cost2, _ := validation.GetPasswordCost(hash2)
|
||||
if cost2 < bcrypt.MinCost {
|
||||
t.Errorf("Cost should be at least MinCost")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckPasswordAndUpgrade(t *testing.T) {
|
||||
password := "testPassword123"
|
||||
|
||||
// 使用低 cost 创建 hash
|
||||
hash, _ := validation.HashPasswordWithCost(password, 4)
|
||||
|
||||
// 检查并尝试升级
|
||||
match, needUpgrade, newHash, err := validation.CheckPasswordAndUpgrade(hash, password, 12)
|
||||
if err != nil {
|
||||
t.Fatalf("CheckPasswordAndUpgrade error: %v", err)
|
||||
}
|
||||
|
||||
if !match {
|
||||
t.Error("Password should match")
|
||||
}
|
||||
|
||||
if !needUpgrade {
|
||||
t.Error("Should need upgrade (cost 4 -> 12)")
|
||||
}
|
||||
|
||||
if newHash == "" {
|
||||
t.Error("New hash should be provided")
|
||||
}
|
||||
|
||||
// 验证新 hash 可用
|
||||
if !validation.CheckPassword(newHash, password) {
|
||||
t.Error("New hash should work")
|
||||
}
|
||||
|
||||
// 使用高 cost,不需要升级
|
||||
hash2, _ := validation.HashPasswordWithCost(password, 12)
|
||||
match2, needUpgrade2, _, _ := validation.CheckPasswordAndUpgrade(hash2, password, 12)
|
||||
if !match2 {
|
||||
t.Error("Password should match")
|
||||
}
|
||||
if needUpgrade2 {
|
||||
t.Error("Should not need upgrade")
|
||||
}
|
||||
|
||||
// 错误密码
|
||||
match3, _, _, _ := validation.CheckPasswordAndUpgrade(hash, "wrongPassword", 12)
|
||||
if match3 {
|
||||
t.Error("Wrong password should not match")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPasswordCost(t *testing.T) {
|
||||
password := "testPassword123"
|
||||
hash, _ := validation.HashPasswordWithCost(password, 12)
|
||||
|
||||
cost, err := validation.GetPasswordCost(hash)
|
||||
if err != nil {
|
||||
t.Fatalf("GetPasswordCost error: %v", err)
|
||||
}
|
||||
if cost != 12 {
|
||||
t.Errorf("Cost = %d, want 12", cost)
|
||||
}
|
||||
|
||||
// 无效 hash
|
||||
_, err = validation.GetPasswordCost("invalidHash")
|
||||
if err == nil {
|
||||
t.Error("GetPasswordCost should fail with invalid hash")
|
||||
}
|
||||
}
|
||||
|
||||
// ===== Validation Errors Tests =====
|
||||
|
||||
func TestValidationErrors(t *testing.T) {
|
||||
errors := validation.ValidationErrors{
|
||||
{Field: "name", Label: "姓名", Message: "必填"},
|
||||
{Field: "email", Label: "邮箱", Message: "格式错误"},
|
||||
}
|
||||
|
||||
// Error 方法
|
||||
errStr := errors.Error()
|
||||
if errStr != "姓名: 必填; 邮箱: 格式错误" {
|
||||
t.Errorf("Error() = %s", errStr)
|
||||
}
|
||||
|
||||
// ToMap 方法
|
||||
m := errors.ToMap()
|
||||
if m["name"] != "必填" {
|
||||
t.Error("ToMap failed")
|
||||
}
|
||||
|
||||
// ToLabelMap 方法
|
||||
lm := errors.ToLabelMap()
|
||||
if lm["姓名"] != "必填" {
|
||||
t.Error("ToLabelMap failed")
|
||||
}
|
||||
|
||||
// First 方法
|
||||
first := errors.First()
|
||||
if first.Field != "name" {
|
||||
t.Error("First failed")
|
||||
}
|
||||
|
||||
// FirstMessage 方法
|
||||
msg := errors.FirstMessage()
|
||||
if msg != "必填" {
|
||||
t.Errorf("FirstMessage = %s", msg)
|
||||
}
|
||||
|
||||
// 空 errors
|
||||
emptyErrors := validation.ValidationErrors{}
|
||||
if emptyErrors.First() != nil {
|
||||
t.Error("Empty First should return nil")
|
||||
}
|
||||
if emptyErrors.FirstMessage() != "" {
|
||||
t.Error("Empty FirstMessage should return empty")
|
||||
}
|
||||
}
|
||||
|
||||
// ===== Struct Validation Tests =====
|
||||
|
||||
type TestUser struct {
|
||||
Name string `json:"name" label:"姓名" binding:"required" msg_required:"姓名不能为空"`
|
||||
Email string `json:"email" label:"邮箱" binding:"required,email" msg_required:"邮箱不能为空" msg_email:"邮箱格式不正确"`
|
||||
Age int `json:"age" label:"年龄" binding:"gte=0,lte=150"`
|
||||
Password string `json:"password" binding:"min=8" msg_min:"密码至少8位"`
|
||||
}
|
||||
|
||||
func TestValidateStruct(t *testing.T) {
|
||||
validation.InitValidator()
|
||||
|
||||
// 有效数据
|
||||
validUser := TestUser{
|
||||
Name: "张三",
|
||||
Email: "test@example.com",
|
||||
Age: 25,
|
||||
Password: "password123",
|
||||
}
|
||||
|
||||
errors := validation.ValidateStruct(validUser)
|
||||
if errors != nil {
|
||||
t.Errorf("Valid struct should have no errors: %v", errors)
|
||||
}
|
||||
|
||||
// 无效数据
|
||||
invalidUser := TestUser{
|
||||
Name: "",
|
||||
Email: "invalid-email",
|
||||
Age: 200,
|
||||
Password: "short",
|
||||
}
|
||||
|
||||
errors2 := validation.ValidateStruct(invalidUser)
|
||||
if errors2 == nil {
|
||||
t.Error("Invalid struct should have errors")
|
||||
}
|
||||
|
||||
// 检查错误数量
|
||||
if len(errors2) < 3 {
|
||||
t.Errorf("Should have at least 3 errors, got %d", len(errors2))
|
||||
}
|
||||
|
||||
// 检查自定义消息
|
||||
firstMsg := errors2.FirstMessage()
|
||||
if firstMsg != "姓名不能为空" {
|
||||
t.Errorf("FirstMessage = %s, want '姓名不能为空'", firstMsg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateStructNil(t *testing.T) {
|
||||
// 空结构体
|
||||
emptyUser := TestUser{}
|
||||
errors := validation.ValidateStruct(emptyUser)
|
||||
if errors == nil {
|
||||
t.Error("Empty struct should have errors")
|
||||
}
|
||||
}
|
||||
|
||||
type TestPhone struct {
|
||||
Phone string `json:"phone" binding:"phone" msg_phone:"手机号格式不正确"`
|
||||
}
|
||||
|
||||
func TestValidatePhone(t *testing.T) {
|
||||
validation.InitValidator()
|
||||
|
||||
// 有效手机号
|
||||
valid := TestPhone{Phone: "13812345678"}
|
||||
errors := validation.ValidateStruct(valid)
|
||||
if errors != nil {
|
||||
t.Errorf("Valid phone should pass: %v", errors)
|
||||
}
|
||||
|
||||
// 无效手机号 - 长度错误
|
||||
invalidLen := TestPhone{Phone: "1234567"}
|
||||
errors2 := validation.ValidateStruct(invalidLen)
|
||||
if errors2 == nil {
|
||||
t.Error("Invalid phone length should fail")
|
||||
}
|
||||
|
||||
// 无效手机号 - 不以1开头
|
||||
invalidPrefix := TestPhone{Phone: "23812345678"}
|
||||
errors3 := validation.ValidateStruct(invalidPrefix)
|
||||
if errors3 == nil {
|
||||
t.Error("Phone not starting with 1 should fail")
|
||||
}
|
||||
}
|
||||
|
||||
type TestUsername struct {
|
||||
Username string `json:"username" binding:"username" msg_username:"用户名格式不正确"`
|
||||
}
|
||||
|
||||
func TestValidateUsername(t *testing.T) {
|
||||
validation.InitValidator()
|
||||
|
||||
tests := []struct {
|
||||
username string
|
||||
valid bool
|
||||
}{
|
||||
{"abc123", true}, // 有效
|
||||
{"Abc123", true}, // 有效(大写开头)
|
||||
{"user_name", true}, // 有效(包含下划线)
|
||||
{"123abc", false}, // 无效(数字开头)
|
||||
{"ab", false}, // 无效(太短)
|
||||
{"a!bc", false}, // 无效(特殊字符)
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
u := TestUsername{Username: tt.username}
|
||||
errors := validation.ValidateStruct(u)
|
||||
valid := errors == nil
|
||||
if valid != tt.valid {
|
||||
t.Errorf("Username %s: valid=%v, want %v", tt.username, valid, tt.valid)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ===== Benchmarks =====
|
||||
|
||||
func BenchmarkHashPassword(b *testing.B) {
|
||||
password := "testPassword123"
|
||||
for i := 0; i < b.N; i++ {
|
||||
validation.HashPassword(password)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkCheckPassword(b *testing.B) {
|
||||
password := "testPassword123"
|
||||
hash, _ := validation.HashPassword(password)
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
validation.CheckPassword(hash, password)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkValidatePassword(b *testing.B) {
|
||||
password := "TestPassword123"
|
||||
for i := 0; i < b.N; i++ {
|
||||
validation.ValidatePassword(password)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,397 @@
|
||||
package validation
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gin-gonic/gin/binding"
|
||||
"github.com/go-playground/validator/v10"
|
||||
)
|
||||
|
||||
// Validator 全局验证器实例
|
||||
var Validator *validator.Validate
|
||||
|
||||
// ValidationError 验证错误
|
||||
type ValidationError struct {
|
||||
Field string `json:"field"` // 字段名(使用 label 或 json tag)
|
||||
Label string `json:"label"` // 字段中文名(用于显示)
|
||||
Message string `json:"message"` // 错误消息
|
||||
}
|
||||
|
||||
// ValidationErrors 验证错误列表
|
||||
type ValidationErrors []ValidationError
|
||||
|
||||
// Error 实现 error 接口
|
||||
func (ve ValidationErrors) Error() string {
|
||||
var msgs []string
|
||||
for _, e := range ve {
|
||||
if e.Label != "" {
|
||||
msgs = append(msgs, e.Label+": "+e.Message)
|
||||
} else {
|
||||
msgs = append(msgs, e.Field+": "+e.Message)
|
||||
}
|
||||
}
|
||||
return strings.Join(msgs, "; ")
|
||||
}
|
||||
|
||||
// ToMap 转换为 map
|
||||
func (ve ValidationErrors) ToMap() map[string]string {
|
||||
m := make(map[string]string)
|
||||
for _, e := range ve {
|
||||
m[e.Field] = e.Message
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// ToLabelMap 转换为带标签的 map
|
||||
func (ve ValidationErrors) ToLabelMap() map[string]string {
|
||||
m := make(map[string]string)
|
||||
for _, e := range ve {
|
||||
if e.Label != "" {
|
||||
m[e.Label] = e.Message
|
||||
} else {
|
||||
m[e.Field] = e.Message
|
||||
}
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// First 获取第一个错误
|
||||
func (ve ValidationErrors) First() *ValidationError {
|
||||
if len(ve) == 0 {
|
||||
return nil
|
||||
}
|
||||
return &ve[0]
|
||||
}
|
||||
|
||||
// FirstMessage 获取第一个错误消息
|
||||
func (ve ValidationErrors) FirstMessage() string {
|
||||
if len(ve) == 0 {
|
||||
return ""
|
||||
}
|
||||
return ve[0].Message
|
||||
}
|
||||
|
||||
// InitValidator 初始化验证器
|
||||
func InitValidator() {
|
||||
if v, ok := binding.Validator.Engine().(*validator.Validate); ok {
|
||||
Validator = v
|
||||
|
||||
// 注册自定义标签名函数(优先使用 label,其次 json)
|
||||
v.RegisterTagNameFunc(func(fld reflect.StructField) string {
|
||||
// 优先使用 label tag 作为字段显示名
|
||||
label := fld.Tag.Get("label")
|
||||
if label != "" {
|
||||
return label
|
||||
}
|
||||
|
||||
// 其次使用 json tag
|
||||
name := strings.SplitN(fld.Tag.Get("json"), ",", 2)[0]
|
||||
if name == "-" {
|
||||
return ""
|
||||
}
|
||||
return name
|
||||
})
|
||||
|
||||
// 注册自定义验证规则
|
||||
registerCustomValidations(v)
|
||||
}
|
||||
}
|
||||
|
||||
// registerCustomValidations 注册自定义验证规则
|
||||
func registerCustomValidations(v *validator.Validate) {
|
||||
// 密码强度验证
|
||||
v.RegisterValidation("password", func(fl validator.FieldLevel) bool {
|
||||
password := fl.Field().String()
|
||||
valid, _ := ValidatePassword(password)
|
||||
return valid
|
||||
})
|
||||
|
||||
// 手机号验证(中国大陆)
|
||||
v.RegisterValidation("phone", func(fl validator.FieldLevel) bool {
|
||||
phone := fl.Field().String()
|
||||
if len(phone) != 11 {
|
||||
return false
|
||||
}
|
||||
return strings.HasPrefix(phone, "1")
|
||||
})
|
||||
|
||||
// 用户名验证(字母开头,允许字母数字下划线)
|
||||
v.RegisterValidation("username", func(fl validator.FieldLevel) bool {
|
||||
username := fl.Field().String()
|
||||
if len(username) < 3 || len(username) > 20 {
|
||||
return false
|
||||
}
|
||||
if !isLetter(rune(username[0])) {
|
||||
return false
|
||||
}
|
||||
for _, r := range username {
|
||||
if !isLetter(r) && !isDigit(r) && r != '_' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
// 手机号严格验证(验证运营商号段)
|
||||
v.RegisterValidation("phone_strict", func(fl validator.FieldLevel) bool {
|
||||
phone := fl.Field().String()
|
||||
if len(phone) != 11 {
|
||||
return false
|
||||
}
|
||||
if !strings.HasPrefix(phone, "1") {
|
||||
return false
|
||||
}
|
||||
// 检查号段
|
||||
prefix := phone[:3]
|
||||
validPrefixes := []string{
|
||||
"130", "131", "132", "133", "134", "135", "136", "137", "138", "139",
|
||||
"145", "146", "147", "148", "149",
|
||||
"150", "151", "152", "153", "155", "156", "157", "158", "159",
|
||||
"166", "167",
|
||||
"170", "171", "172", "173", "174", "175", "176", "177", "178",
|
||||
"180", "181", "182", "183", "184", "185", "186", "187", "188", "189",
|
||||
"191", "198", "199",
|
||||
}
|
||||
for _, p := range validPrefixes {
|
||||
if prefix == p {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
})
|
||||
|
||||
// 身份证号验证(简化版)
|
||||
v.RegisterValidation("idcard", func(fl validator.FieldLevel) bool {
|
||||
id := fl.Field().String()
|
||||
if len(id) != 18 && len(id) != 15 {
|
||||
return false
|
||||
}
|
||||
// 简化验证:只检查长度和基本格式
|
||||
for i, c := range id {
|
||||
if i == len(id)-1 && len(id) == 18 {
|
||||
// 最后一位可以是 X
|
||||
if !isDigit(c) && c != 'X' && c != 'x' {
|
||||
return false
|
||||
}
|
||||
} else {
|
||||
if !isDigit(c) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
func isLetter(r rune) bool {
|
||||
return (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z')
|
||||
}
|
||||
|
||||
func isDigit(r rune) bool {
|
||||
return r >= '0' && r <= '9'
|
||||
}
|
||||
|
||||
// ValidateStruct 验证结构体
|
||||
func ValidateStruct(s any) ValidationErrors {
|
||||
if Validator == nil {
|
||||
InitValidator()
|
||||
}
|
||||
|
||||
err := Validator.Struct(s)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return parseValidationErrors(err, s)
|
||||
}
|
||||
|
||||
// parseValidationErrors 解析验证错误(支持自定义错误消息)
|
||||
func parseValidationErrors(err error, s any) ValidationErrors {
|
||||
var errors ValidationErrors
|
||||
|
||||
if validationErrors, ok := err.(validator.ValidationErrors); ok {
|
||||
for _, e := range validationErrors {
|
||||
fieldName := e.Field()
|
||||
label := fieldName // Field() 返回的是 label 或 json tag
|
||||
|
||||
// 尝试获取原始字段名和自定义错误消息
|
||||
if s != nil {
|
||||
t := reflect.TypeOf(s)
|
||||
if t.Kind() == reflect.Ptr {
|
||||
t = t.Elem()
|
||||
}
|
||||
if t.Kind() == reflect.Struct {
|
||||
// 获取原始字段名
|
||||
originalField := getOriginalFieldName(t, e.StructField())
|
||||
if originalField != "" {
|
||||
fieldName = originalField
|
||||
}
|
||||
|
||||
// 获取自定义错误消息
|
||||
field, found := t.FieldByName(e.StructField())
|
||||
if found {
|
||||
customMsg := getCustomErrorMessage(field, e.Tag())
|
||||
if customMsg != "" {
|
||||
errors = append(errors, ValidationError{
|
||||
Field: fieldName,
|
||||
Label: label,
|
||||
Message: customMsg,
|
||||
})
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
errors = append(errors, ValidationError{
|
||||
Field: fieldName,
|
||||
Label: label,
|
||||
Message: getErrorMessage(e),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return errors
|
||||
}
|
||||
|
||||
// getOriginalFieldName 获取原始字段名(从 json tag)
|
||||
func getOriginalFieldName(t reflect.Type, structField string) string {
|
||||
field, found := t.FieldByName(structField)
|
||||
if !found {
|
||||
return ""
|
||||
}
|
||||
|
||||
jsonTag := field.Tag.Get("json")
|
||||
if jsonTag == "" || jsonTag == "-" {
|
||||
return structField
|
||||
}
|
||||
|
||||
name := strings.SplitN(jsonTag, ",", 2)[0]
|
||||
if name == "" {
|
||||
return structField
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
// getCustomErrorMessage 获取自定义错误消息
|
||||
// 支持格式:
|
||||
// - error:"自定义错误消息"
|
||||
// - msg_required:"必填项" (针对特定验证规则)
|
||||
// - msg_min:"最少5个字符"
|
||||
func getCustomErrorMessage(field reflect.StructField, tag string) string {
|
||||
// 优先查找特定规则的错误消息
|
||||
specificTag := fmt.Sprintf("msg_%s", tag)
|
||||
msg := field.Tag.Get(specificTag)
|
||||
if msg != "" {
|
||||
return msg
|
||||
}
|
||||
|
||||
// 其次查找通用错误消息
|
||||
msg = field.Tag.Get("error")
|
||||
if msg != "" {
|
||||
return msg
|
||||
}
|
||||
|
||||
// 最后查找 msg tag
|
||||
msg = field.Tag.Get("msg")
|
||||
if msg != "" {
|
||||
return msg
|
||||
}
|
||||
|
||||
return ""
|
||||
}
|
||||
|
||||
// getErrorMessage 获取默认验证错误消息
|
||||
func getErrorMessage(e validator.FieldError) string {
|
||||
switch e.Tag() {
|
||||
case "required":
|
||||
return "此字段为必填项"
|
||||
case "email":
|
||||
return "邮箱格式不正确"
|
||||
case "min":
|
||||
return fmt.Sprintf("长度不能少于 %s 个字符", e.Param())
|
||||
case "max":
|
||||
return fmt.Sprintf("长度不能超过 %s 个字符", e.Param())
|
||||
case "len":
|
||||
return fmt.Sprintf("长度必须为 %s 个字符", e.Param())
|
||||
case "gte":
|
||||
return fmt.Sprintf("必须大于或等于 %s", e.Param())
|
||||
case "lte":
|
||||
return fmt.Sprintf("必须小于或等于 %s", e.Param())
|
||||
case "gt":
|
||||
return fmt.Sprintf("必须大于 %s", e.Param())
|
||||
case "lt":
|
||||
return fmt.Sprintf("必须小于 %s", e.Param())
|
||||
case "eq":
|
||||
return fmt.Sprintf("必须等于 %s", e.Param())
|
||||
case "ne":
|
||||
return fmt.Sprintf("不能等于 %s", e.Param())
|
||||
case "oneof":
|
||||
return fmt.Sprintf("必须是以下值之一: %s", e.Param())
|
||||
case "url":
|
||||
return "URL 格式不正确"
|
||||
case "uri":
|
||||
return "URI 格式不正确"
|
||||
case "uuid":
|
||||
return "UUID 格式不正确"
|
||||
case "alphanum":
|
||||
return "只能包含字母和数字"
|
||||
case "alpha":
|
||||
return "只能包含字母"
|
||||
case "numeric":
|
||||
return "必须是数字"
|
||||
case "password":
|
||||
return "密码强度不足,需包含大小写字母和数字,至少8位"
|
||||
case "phone":
|
||||
return "手机号格式不正确"
|
||||
case "phone_strict":
|
||||
return "手机号无效,请输入正确的手机号"
|
||||
case "username":
|
||||
return "用户名必须以字母开头,只能包含字母、数字和下划线,长度3-20"
|
||||
case "idcard":
|
||||
return "身份证号格式不正确"
|
||||
default:
|
||||
return fmt.Sprintf("验证失败: %s", e.Tag())
|
||||
}
|
||||
}
|
||||
|
||||
// BindAndValidate 绑定并验证请求
|
||||
func BindAndValidate(c *gin.Context, req any) ValidationErrors {
|
||||
if err := c.ShouldBind(req); err != nil {
|
||||
return parseValidationErrors(err, req)
|
||||
}
|
||||
return ValidateStruct(req)
|
||||
}
|
||||
|
||||
// ShouldBindAndValidate 绑定并验证请求,返回是否成功
|
||||
func ShouldBindAndValidate(c *gin.Context, req any) (ValidationErrors, bool) {
|
||||
errors := BindAndValidate(c, req)
|
||||
return errors, len(errors) == 0
|
||||
}
|
||||
|
||||
// BindJSON 绑定 JSON 并验证
|
||||
func BindJSON(c *gin.Context, req any) ValidationErrors {
|
||||
if err := c.ShouldBindJSON(req); err != nil {
|
||||
return parseValidationErrors(err, req)
|
||||
}
|
||||
return ValidateStruct(req)
|
||||
}
|
||||
|
||||
// BindQuery 绑定 Query 并验证
|
||||
func BindQuery(c *gin.Context, req any) ValidationErrors {
|
||||
if err := c.ShouldBindQuery(req); err != nil {
|
||||
return parseValidationErrors(err, req)
|
||||
}
|
||||
return ValidateStruct(req)
|
||||
}
|
||||
|
||||
// BindForm 绑定 Form 并验证
|
||||
func BindForm(c *gin.Context, req any) ValidationErrors {
|
||||
if err := c.ShouldBind(req); err != nil {
|
||||
return parseValidationErrors(err, req)
|
||||
}
|
||||
return ValidateStruct(req)
|
||||
}
|
||||
Reference in New Issue
Block a user