u
This commit is contained in:
1 parent
69d03709ba
commit
44fb48ef6a
9 files changed
+625
-25
No files matched your search
@@ -4,8 +4,10 @@ go 1.25.0
|
||||
|
||||
require (
|
||||
github.com/golang-jwt/jwt/v5 v5.3.0
|
||||
github.com/jmoiron/sqlx v1.4.0
|
||||
github.com/lib/pq v1.10.9
|
||||
github.com/spf13/viper v1.21.0
|
||||
github.com/valyala/fasthttp v1.68.0
|
||||
go.uber.org/zap v1.27.0
|
||||
)
|
||||
|
||||
@@ -13,7 +15,6 @@ require (
|
||||
github.com/andybalholm/brotli v1.2.0 // indirect
|
||||
github.com/fsnotify/fsnotify v1.9.0 // indirect
|
||||
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
|
||||
github.com/jmoiron/sqlx v1.4.0 // indirect
|
||||
github.com/klauspost/compress v1.18.1 // indirect
|
||||
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
|
||||
github.com/sagikazarmark/locafero v0.12.0 // indirect
|
||||
@@ -23,7 +24,6 @@ require (
|
||||
github.com/spf13/pflag v1.0.10 // indirect
|
||||
github.com/subosito/gotenv v1.6.0 // indirect
|
||||
github.com/valyala/bytebufferpool v1.0.0 // indirect
|
||||
github.com/valyala/fasthttp v1.68.0 // indirect
|
||||
go.uber.org/multierr v1.11.0 // indirect
|
||||
go.yaml.in/yaml/v3 v3.0.4 // indirect
|
||||
golang.org/x/sys v0.37.0 // indirect
|
||||
|
||||
@@ -1,13 +1,17 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"base-go-v2/internal/db"
|
||||
response "base-go-v2/internal/reponse"
|
||||
"base-go-v2/internal/router"
|
||||
"base-go-v2/internal/utils/mapx"
|
||||
"base-go-v2/internal/utils/uid"
|
||||
)
|
||||
|
||||
func InitAuthRouter(r *router.Router) {
|
||||
group := r.Group("auth")
|
||||
group.POST("/login/default", loginDefault)
|
||||
group.POST("/register/default", registerDefault)
|
||||
}
|
||||
|
||||
// 默认登录
|
||||
@@ -15,13 +19,41 @@ func loginDefault(c *router.Context) error {
|
||||
return response.Success(c).Send()
|
||||
}
|
||||
|
||||
// 注册
|
||||
func registerDefault(c *router.Context) error {
|
||||
var req struct {
|
||||
Email string `json:"email"`
|
||||
Password string `json:"password"`
|
||||
params, err := c.GetBodyWithRequired("account", "password")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
account := params.GetString("account")
|
||||
password := params.GetString("password")
|
||||
|
||||
// 构造查询条件,检查账号是否已存在
|
||||
accountQuery := mapx.New().SetKV("account", account)
|
||||
|
||||
existingUsers, err := db.Find("user_info", accountQuery)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := c.BindJSON(&req); err != nil {
|
||||
if len(existingUsers) > 0 {
|
||||
return response.Error(c).Message("该账号已被注册").Send()
|
||||
}
|
||||
|
||||
// 构造新用户数据
|
||||
newUserData := mapx.New().SetKV(
|
||||
"id", uid.NextID(),
|
||||
"account", account,
|
||||
"password", password,
|
||||
)
|
||||
|
||||
// 插入新用户
|
||||
insertedRows, err := db.InsertOne("user_info", newUserData)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if insertedRows != 1 {
|
||||
return response.Error(c).Message("注册失败").Send()
|
||||
}
|
||||
|
||||
return response.Success(c).Message("注册成功").Send()
|
||||
}
|
||||
@@ -1,29 +1,45 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"base-go-v2/internal/errs"
|
||||
"base-go-v2/internal/logx"
|
||||
"base-go-v2/internal/reponse"
|
||||
"base-go-v2/internal/router"
|
||||
"errors"
|
||||
"runtime/debug"
|
||||
)
|
||||
|
||||
func ErrorMiddleware() router.HandlerFunc {
|
||||
// 可扩展的业务错误类型列表
|
||||
businessErrors := []interface{}{
|
||||
(*errs.ValidationError)(nil),
|
||||
// (*errs.BusinessError)(nil),
|
||||
// (*errs.AuthError)(nil),
|
||||
}
|
||||
|
||||
return func(c *router.Context) error {
|
||||
err := c.Next()
|
||||
|
||||
if err != nil {
|
||||
stack := string(debug.Stack())
|
||||
|
||||
logx.Logger.Error("请求错误",
|
||||
logx.String("path", string(c.RequestCtx.Path())),
|
||||
logx.String("method", string(c.RequestCtx.Method())),
|
||||
logx.String("error", err.Error()),
|
||||
logx.String("stack", stack),
|
||||
)
|
||||
|
||||
return response.Error(c).Message(err.Error()).Send()
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return nil
|
||||
// 判断是否属于业务错误
|
||||
for _, be := range businessErrors {
|
||||
// be 是 nil 指针类型
|
||||
if errors.As(err, &be) {
|
||||
// 业务错误,不记录日志
|
||||
return response.Error(c).Message(err.Error()).Send()
|
||||
}
|
||||
}
|
||||
|
||||
// 系统错误,记录日志
|
||||
stack := string(debug.Stack())
|
||||
logx.Logger.Error("请求错误",
|
||||
logx.String("path", string(c.RequestCtx.Path())),
|
||||
logx.String("method", string(c.RequestCtx.Method())),
|
||||
logx.String("error", err.Error()),
|
||||
logx.String("stack", stack),
|
||||
)
|
||||
return response.Error(c).Message("服务器内部错误").Send()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package errs
|
||||
|
||||
import "fmt"
|
||||
|
||||
type ValidationError struct {
|
||||
field string
|
||||
message string
|
||||
}
|
||||
|
||||
func (e *ValidationError) Error() string {
|
||||
return fmt.Sprintf("字段 %s: %s", e.field, e.message)
|
||||
}
|
||||
|
||||
func NewValidationError(field, message string) error {
|
||||
return &ValidationError{
|
||||
field: field,
|
||||
message: message,
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,8 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"base-go-v2/internal/utils/mapx"
|
||||
"base-go-v2/internal/utils/validate"
|
||||
"encoding/json"
|
||||
"sync"
|
||||
|
||||
@@ -72,6 +74,37 @@ func (c *Context) BindJSON(obj interface{}) error {
|
||||
return json.Unmarshal(body, obj)
|
||||
}
|
||||
|
||||
// BodyMap 将请求体解析为 mapx.M
|
||||
func (c *Context) BodyMap() (mapx.M, error) {
|
||||
body, err := c.Body() // 使用 BodyCache 读取请求体
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
m := mapx.New()
|
||||
if err := json.Unmarshal(body, &m); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (c *Context) GetBodyWithRequired(keys ...string) (mapx.M, error) {
|
||||
// 解析请求体
|
||||
req, err := c.BodyMap()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 校验必填字段
|
||||
if err := validate.ValidateNotEmpty(req, keys...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 返回整个请求体(包含所有字段)
|
||||
return req, nil
|
||||
}
|
||||
|
||||
// PostForm 获取表单参数
|
||||
func (c *Context) PostForm(key string) string {
|
||||
return string(c.RequestCtx.FormValue(key))
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
package mapx
|
||||
|
||||
// M 是 map[string]interface{} 的别名
|
||||
type M map[string]interface{}
|
||||
|
||||
// New 创建一个空 M
|
||||
func New() M {
|
||||
return M{}
|
||||
}
|
||||
|
||||
// Set 设置单个键值,支持链式调用
|
||||
func (m M) Set(key string, value interface{}) M {
|
||||
m[key] = value
|
||||
return m
|
||||
}
|
||||
|
||||
// SetKV 批量设置键值,可变参数形式
|
||||
func (m M) SetKV(kvs ...interface{}) M {
|
||||
if len(kvs)%2 != 0 {
|
||||
panic("SetKV 参数必须成对出现")
|
||||
}
|
||||
for i := 0; i < len(kvs); i += 2 {
|
||||
key, ok := kvs[i].(string)
|
||||
if !ok {
|
||||
panic("SetKV 键必须是字符串")
|
||||
}
|
||||
m[key] = kvs[i+1]
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// Get 获取值
|
||||
func (m M) Get(key string) interface{} {
|
||||
return m[key]
|
||||
}
|
||||
|
||||
// GetString 获取字符串类型
|
||||
func (m M) GetString(key string) string {
|
||||
if v, ok := m[key]; ok {
|
||||
if s, ok := v.(string); ok {
|
||||
return s
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// GetInt 获取整数类型
|
||||
func (m M) GetInt(key string) int {
|
||||
if v, ok := m[key]; ok {
|
||||
switch val := v.(type) {
|
||||
case int:
|
||||
return val
|
||||
case int8:
|
||||
return int(val)
|
||||
case int16:
|
||||
return int(val)
|
||||
case int32:
|
||||
return int(val)
|
||||
case int64:
|
||||
return int(val)
|
||||
case float32:
|
||||
return int(val)
|
||||
case float64:
|
||||
return int(val)
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// GetFloat 获取浮点数类型
|
||||
func (m M) GetFloat(key string) float64 {
|
||||
if v, ok := m[key]; ok {
|
||||
switch val := v.(type) {
|
||||
case float32:
|
||||
return float64(val)
|
||||
case float64:
|
||||
return val
|
||||
case int:
|
||||
return float64(val)
|
||||
case int8:
|
||||
return float64(val)
|
||||
case int16:
|
||||
return float64(val)
|
||||
case int32:
|
||||
return float64(val)
|
||||
case int64:
|
||||
return float64(val)
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// GetBool 获取布尔类型
|
||||
func (m M) GetBool(key string) bool {
|
||||
if v, ok := m[key]; ok {
|
||||
if b, ok := v.(bool); ok {
|
||||
return b
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Merge 合并另一个 M,后者会覆盖前者同名字段
|
||||
func (m M) Merge(other M) M {
|
||||
for k, v := range other {
|
||||
m[k] = v
|
||||
}
|
||||
return m
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package validate
|
||||
|
||||
import (
|
||||
"base-go-v2/internal/errs"
|
||||
"errors"
|
||||
"reflect"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ValidateNotEmpty 校验结构体中指定字段不能为空
|
||||
func ValidateNotEmpty(s interface{}, fields ...string) error {
|
||||
v := reflect.ValueOf(s)
|
||||
if v.Kind() == reflect.Ptr {
|
||||
v = v.Elem()
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
Reference in new issue
Block a user