This commit is contained in:
oneao committed 2025-11-04 22:39:02 +08:00
1 parent 69d03709ba
commit 44fb48ef6a
9 files changed
+625 -25

No files matched your search

+2 -2
View File
@@ -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
}