chore: initial commit

This commit is contained in:
2026-07-22 17:50:33 +08:00
parent 82936a8987
commit ee9f50b859
107 changed files with 8346 additions and 0 deletions

101
internal/middleware/auth.go Normal file
View File

@@ -0,0 +1,101 @@
package middleware
import (
"context"
"net/http"
db "server/internal/db/sqlc"
"server/internal/pkg/cache"
"server/internal/pkg/errs"
"server/internal/pkg/httputil"
"strings"
"github.com/go-chi/chi/v5"
)
const (
IsAdminKey contextKey = "is_admin"
)
type AuthMiddleware struct {
queries *db.Queries
cache *cache.Caches
}
func NewAuthMiddleware(queries *db.Queries, cache *cache.Caches) *AuthMiddleware {
return &AuthMiddleware{
cache: cache,
queries: queries,
}
}
func (m *AuthMiddleware) Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
claims, ok := GetClaims(ctx)
if !ok || claims.UserID == 0 {
httputil.Fail(w, errs.ErrUnauthorized)
return
}
// 判断是否有管理员权限 目前只判断uid是否为1
isAdmin := userIsAdmin(claims.UserID)
if isAdmin {
ctx = context.WithValue(ctx, IsAdminKey, isAdmin)
next.ServeHTTP(w, r.WithContext(ctx))
return
}
hasPermission, err := userHasApiPermission(ctx, r, m.queries, claims.UserID, m.cache)
if err != nil {
httputil.Fail(w, err)
return
}
if !hasPermission {
httputil.Fail(w, errs.ErrPermissionDenied)
return
}
next.ServeHTTP(w, r)
})
}
func userIsAdmin(uid int32) bool {
if uid == 1 {
return true
}
return false
}
func userHasApiPermission(ctx context.Context, r *http.Request, queries *db.Queries, uid int32, cache *cache.Caches) (bool, error) {
var (
apis []db.GetSysUserApisRow
err error
)
apis, ok := cache.SysUserApisCache.GetIfPresent(uid)
if !ok {
apis, err = queries.GetSysUserApis(ctx, uid)
if err != nil {
return false, err
}
cache.SysUserApisCache.Set(uid, apis)
}
requestPath := chi.RouteContext(r.Context()).RoutePattern()
requestPath = strings.TrimPrefix(requestPath, "/api")
requestMethod := r.Method
for _, api := range apis {
if api.Path == requestPath && api.Method == requestMethod {
return true, nil
}
}
return false, nil
}

184
internal/middleware/jwt.go Normal file
View File

@@ -0,0 +1,184 @@
package middleware
import (
"context"
"errors"
"net/http"
"server/internal/config"
"server/internal/pkg/errs"
"server/internal/pkg/httputil"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
)
type JWTMiddleware struct {
cfg *config.JWTConfig
}
type contextKey string
const UserContextKey contextKey = "user"
const RefreshTokenType = "refresh"
type Claims struct {
UserID int32 `json:"user_id"`
jwt.RegisteredClaims
}
type RefreshClaims struct {
UserID int32 `json:"user_id"`
Type string `json:"type"`
jwt.RegisteredClaims
}
func NewJWTMiddleware(cfg *config.Config) *JWTMiddleware {
jwtCfg := cfg.JWTConfig
return &JWTMiddleware{cfg: &jwtCfg}
}
func GetClaims(ctx context.Context) (*Claims, bool) {
claims, ok := ctx.Value(UserContextKey).(*Claims)
return claims, ok
}
// ParseToken 解析accessToken
func (m *JWTMiddleware) ParseToken(tokenStr string) (*Claims, error) {
token, err := jwt.ParseWithClaims(
tokenStr,
&Claims{},
func(token *jwt.Token) (interface{}, error) {
return []byte(m.cfg.Secret), nil
},
)
if err != nil || !token.Valid {
return nil, errs.ErrInvalidToken
}
claims, ok := token.Claims.(*Claims)
if !ok {
return nil, errs.ErrInvalidTokenClaims
}
if claims.UserID == 0 {
return nil, errs.ErrInvalidTokenClaims
}
return claims, nil
}
// ParseRefreshToken 解析RefreshToken
func (m *JWTMiddleware) ParseRefreshToken(tokenStr string) (*RefreshClaims, error) {
token, err := jwt.ParseWithClaims(
tokenStr,
&RefreshClaims{},
func(token *jwt.Token) (interface{}, error) {
return []byte(m.cfg.Secret), nil
},
)
if err != nil {
// Token 已过期
if errors.Is(err, jwt.ErrTokenExpired) {
return nil, errs.ErrExpiredRefreshToken
}
// 签名错误、格式错误、非法 Token
return nil, errs.ErrInvalidRefreshToken
}
if !token.Valid {
return nil, errs.ErrInvalidRefreshToken
}
claims, ok := token.Claims.(*RefreshClaims)
if !ok {
return nil, errs.ErrInvalidRefreshToken
}
if claims.UserID == 0 {
return nil, errs.ErrInvalidRefreshToken
}
if claims.Type != RefreshTokenType {
return nil, errs.ErrInvalidRefreshToken
}
return claims, nil
}
func (m *JWTMiddleware) Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
auth := r.Header.Get("Authorization")
if auth == "" {
httputil.Fail(w, errs.ErrUnauthenticated)
return
}
parts := strings.SplitN(auth, " ", 2)
if len(parts) != 2 || parts[0] != "Bearer" {
httputil.Fail(w, errs.ErrUnauthenticated)
return
}
tokenStr := parts[1]
claims, err := m.ParseToken(tokenStr)
if err != nil {
httputil.Fail(w, errs.ErrInvalidToken)
return
}
ctx := context.WithValue(r.Context(), UserContextKey, claims)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
func (m *JWTMiddleware) GenerateAccessToken(userID int32) (string, time.Time, error) {
now := time.Now()
expiresAt := now.Add(m.cfg.Expire)
accessClaims := Claims{
UserID: userID,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(expiresAt),
IssuedAt: jwt.NewNumericDate(now),
},
}
accessToken := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims)
token, err := accessToken.SignedString([]byte(m.cfg.Secret))
if err != nil {
return "", time.Time{}, err
}
return token, expiresAt, nil
}
func (m *JWTMiddleware) GenerateRefreshToken(userID int32) (string, time.Time, error) {
now := time.Now()
expiresAt := now.Add(m.cfg.RefreshExpire)
refreshClaims := RefreshClaims{
UserID: userID,
Type: RefreshTokenType,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(expiresAt),
IssuedAt: jwt.NewNumericDate(now),
},
}
refreshToken := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims)
token, err := refreshToken.SignedString([]byte(m.cfg.Secret))
if err != nil {
return "", time.Time{}, err
}
return token, expiresAt, nil
}

View File

@@ -0,0 +1,103 @@
package middleware
import (
"encoding/json"
"log/slog"
"net/http"
"server/internal/model/common"
"server/internal/utils"
"time"
gonanoid "github.com/matoous/go-nanoid/v2"
)
type LoggerMiddleware struct {
}
type responseWriter struct {
http.ResponseWriter
statusCode int
bytes int
errorMsg string
}
func (rw *responseWriter) WriteHeader(code int) {
rw.statusCode = code
rw.ResponseWriter.WriteHeader(code)
}
func generateRequestID() string {
id, _ := gonanoid.New(16) // 16字符
return id
}
// extractErrorMessage 从响应体中提取错误信息
func (rw *responseWriter) extractErrorMessage(body []byte) string {
var resp common.Response
if err := json.Unmarshal(body, &resp); err == nil && resp.Message != "" {
return resp.Message
}
return string(body)
}
func (rw *responseWriter) Write(b []byte) (int, error) {
n, err := rw.ResponseWriter.Write(b)
rw.bytes += n
// 只在错误状态码且未记录错误时处理
if rw.statusCode >= 400 && rw.errorMsg == "" && n > 0 {
rw.errorMsg = rw.extractErrorMessage(b[:n])
}
return n, err
}
func NewLoggerMiddleware() *LoggerMiddleware {
return &LoggerMiddleware{}
}
func (m *LoggerMiddleware) Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
reqID := generateRequestID()
wrapped := &responseWriter{
ResponseWriter: w,
statusCode: http.StatusOK,
}
next.ServeHTTP(wrapped, r)
duration := time.Since(start)
ip := utils.ClientIP(r)
fullPath := r.URL.Path
if r.URL.RawQuery != "" {
fullPath = fullPath + "?" + r.URL.RawQuery
}
fields := []any{
"request_id", reqID,
"method", r.Method,
"path", fullPath,
"status", wrapped.statusCode,
"duration", duration,
"client_ip", ip,
}
if wrapped.errorMsg != "" {
fields = append(fields, "error", wrapped.errorMsg)
}
switch {
case wrapped.statusCode >= 500:
slog.Error("request", fields...)
case wrapped.statusCode >= 400:
slog.Warn("request", fields...)
default:
slog.Info("request", fields...)
}
})
}