chore: initial commit of ai site platform
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
102
gateway/internal/jwtmw/jwt.go
Normal file
102
gateway/internal/jwtmw/jwt.go
Normal file
@@ -0,0 +1,102 @@
|
||||
package jwtmw
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/golang-jwt/jwt/v4"
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
TenantID int64 `json:"tenant_id"`
|
||||
UserID int64 `json:"user_id"`
|
||||
Role string `json:"role"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// Middleware 网关侧预检 JWT;公开路径放行。中台仍会再次校验。
|
||||
func Middleware(secret string, publicPrefixes []string) func(http.HandlerFunc) http.HandlerFunc {
|
||||
parser := jwt.NewParser(jwt.WithValidMethods([]string{"HS256"}))
|
||||
return func(next http.HandlerFunc) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
path := r.URL.Path
|
||||
for _, p := range publicPrefixes {
|
||||
if path == p || strings.HasPrefix(path, p) {
|
||||
next(w, r)
|
||||
return
|
||||
}
|
||||
}
|
||||
if r.Method == http.MethodOptions {
|
||||
next(w, r)
|
||||
return
|
||||
}
|
||||
auth := r.Header.Get("Authorization")
|
||||
if !strings.HasPrefix(strings.ToLower(auth), "bearer ") {
|
||||
writeErr(w, http.StatusUnauthorized, "gateway: missing bearer token")
|
||||
return
|
||||
}
|
||||
raw := strings.TrimSpace(auth[len("Bearer "):])
|
||||
if raw == "" || strings.Count(raw, ".") != 2 {
|
||||
writeErr(w, http.StatusUnauthorized, "gateway: malformed token")
|
||||
return
|
||||
}
|
||||
token, err := parser.ParseWithClaims(raw, &Claims{}, func(t *jwt.Token) (any, error) {
|
||||
return []byte(secret), nil
|
||||
})
|
||||
if err != nil || token == nil || !token.Valid {
|
||||
writeErr(w, http.StatusUnauthorized, "gateway: invalid token ("+shortJWTErr(err)+"),请重新登录")
|
||||
return
|
||||
}
|
||||
next(w, r)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func shortJWTErr(err error) string {
|
||||
if err == nil {
|
||||
return "rejected"
|
||||
}
|
||||
msg := err.Error()
|
||||
switch {
|
||||
case strings.Contains(msg, "expired"):
|
||||
return "expired"
|
||||
case strings.Contains(msg, "signature"):
|
||||
return "bad signature"
|
||||
case strings.Contains(msg, "malformed"):
|
||||
return "malformed"
|
||||
case strings.Contains(msg, "used before"):
|
||||
return "not yet valid"
|
||||
default:
|
||||
if len(msg) > 80 {
|
||||
return msg[:80]
|
||||
}
|
||||
return msg
|
||||
}
|
||||
}
|
||||
|
||||
func writeErr(w http.ResponseWriter, code int, msg string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(code)
|
||||
_, _ = w.Write([]byte(`{"code":` + itoa(code) + `,"message":"` + escapeJSON(msg) + `"}`))
|
||||
}
|
||||
|
||||
func escapeJSON(s string) string {
|
||||
s = strings.ReplaceAll(s, `\`, `\\`)
|
||||
s = strings.ReplaceAll(s, `"`, `\"`)
|
||||
s = strings.ReplaceAll(s, "\n", " ")
|
||||
return s
|
||||
}
|
||||
|
||||
func itoa(n int) string {
|
||||
if n == 0 {
|
||||
return "0"
|
||||
}
|
||||
var b [12]byte
|
||||
i := len(b)
|
||||
for n > 0 {
|
||||
i--
|
||||
b[i] = byte('0' + n%10)
|
||||
n /= 10
|
||||
}
|
||||
return string(b[i:])
|
||||
}
|
||||
Reference in New Issue
Block a user