功能更新
This commit is contained in:
23
internal/commonservice/ctx.go
Normal file
23
internal/commonservice/ctx.go
Normal file
@@ -0,0 +1,23 @@
|
||||
package commonservice
|
||||
|
||||
import "github.com/gin-gonic/gin"
|
||||
|
||||
// Gin context 中由 RequireJWT 写入的键名。
|
||||
const (
|
||||
CtxUserID = "user_id"
|
||||
CtxUsername = "username"
|
||||
)
|
||||
|
||||
// UserID 从 gin 上下文读取当前登录用户 ID。
|
||||
func UserID(c *gin.Context) int64 {
|
||||
v, _ := c.Get(CtxUserID)
|
||||
id, _ := v.(int64)
|
||||
return id
|
||||
}
|
||||
|
||||
// Username 从 gin 上下文读取当前登录用户名。
|
||||
func Username(c *gin.Context) string {
|
||||
v, _ := c.Get(CtxUsername)
|
||||
s, _ := v.(string)
|
||||
return s
|
||||
}
|
||||
32
internal/commonservice/errors.go
Normal file
32
internal/commonservice/errors.go
Normal file
@@ -0,0 +1,32 @@
|
||||
package commonservice
|
||||
|
||||
// AppError 是带稳定错误码与 HTTP 状态的业务错误,控制器原样返回 {"error":"CODE"}。
|
||||
type AppError struct {
|
||||
Code string
|
||||
Status int
|
||||
}
|
||||
|
||||
func (e *AppError) Error() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
return e.Code
|
||||
}
|
||||
|
||||
// AsAppError 从 error 中取出 AppError;非业务错误返回 false。
|
||||
func AsAppError(err error) (*AppError, bool) {
|
||||
if err == nil {
|
||||
return nil, false
|
||||
}
|
||||
if ae, ok := err.(*AppError); ok {
|
||||
return ae, true
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func BadRequest(code string) *AppError { return &AppError{Code: code, Status: 400} }
|
||||
func Unauthorized(code string) *AppError { return &AppError{Code: code, Status: 401} }
|
||||
func Forbidden(code string) *AppError { return &AppError{Code: code, Status: 403} }
|
||||
func NotFound(code string) *AppError { return &AppError{Code: code, Status: 404} }
|
||||
func Conflict(code string) *AppError { return &AppError{Code: code, Status: 409} }
|
||||
func Internal(code string) *AppError { return &AppError{Code: code, Status: 500} }
|
||||
67
internal/commonservice/jwt.go
Normal file
67
internal/commonservice/jwt.go
Normal file
@@ -0,0 +1,67 @@
|
||||
package commonservice
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// Claims 是 access / refresh token 的载荷。
|
||||
type Claims struct {
|
||||
UserID int64 `json:"userId"`
|
||||
Username string `json:"username"`
|
||||
Type string `json:"type"` // access | refresh
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// IssueAccess 签发 access token。
|
||||
func IssueAccess(secret string, userID int64, username string, ttlHours int) (string, error) {
|
||||
if ttlHours <= 0 {
|
||||
ttlHours = 2
|
||||
}
|
||||
return issue(secret, userID, username, "access", time.Duration(ttlHours)*time.Hour)
|
||||
}
|
||||
|
||||
// IssueRefresh 签发 refresh token。
|
||||
func IssueRefresh(secret string, userID int64, username string, ttlDays int) (string, error) {
|
||||
if ttlDays <= 0 {
|
||||
ttlDays = 30
|
||||
}
|
||||
return issue(secret, userID, username, "refresh", time.Duration(ttlDays)*24*time.Hour)
|
||||
}
|
||||
|
||||
func issue(secret string, userID int64, username, typ string, ttl time.Duration) (string, error) {
|
||||
now := time.Now()
|
||||
claims := Claims{
|
||||
UserID: userID,
|
||||
Username: username,
|
||||
Type: typ,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
ExpiresAt: jwt.NewNumericDate(now.Add(ttl)),
|
||||
},
|
||||
}
|
||||
t := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return t.SignedString([]byte(secret))
|
||||
}
|
||||
|
||||
// Parse 校验并解析 JWT;typ 非空时要求 claims.Type 匹配。
|
||||
func Parse(secret, token, typ string) (*Claims, error) {
|
||||
parsed, err := jwt.ParseWithClaims(token, &Claims{}, func(t *jwt.Token) (any, error) {
|
||||
if t.Method != jwt.SigningMethodHS256 {
|
||||
return nil, Unauthorized("UNAUTHORIZED")
|
||||
}
|
||||
return []byte(secret), nil
|
||||
})
|
||||
if err != nil || !parsed.Valid {
|
||||
return nil, Unauthorized("UNAUTHORIZED")
|
||||
}
|
||||
claims, ok := parsed.Claims.(*Claims)
|
||||
if !ok || claims.UserID <= 0 {
|
||||
return nil, Unauthorized("UNAUTHORIZED")
|
||||
}
|
||||
if typ != "" && claims.Type != typ {
|
||||
return nil, Unauthorized("UNAUTHORIZED")
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
55
internal/commonservice/team.go
Normal file
55
internal/commonservice/team.go
Normal file
@@ -0,0 +1,55 @@
|
||||
package commonservice
|
||||
|
||||
import "gorm.io/gorm"
|
||||
|
||||
// TeamRole 查询用户在团队中的角色;非成员返回空串。
|
||||
func TeamRole(db *gorm.DB, teamID, userID int64) (string, error) {
|
||||
if teamID <= 0 || userID <= 0 {
|
||||
return "", nil
|
||||
}
|
||||
var role string
|
||||
err := db.Table("team_members").Select("role").
|
||||
Where("team_id = ? AND user_id = ?", teamID, userID).
|
||||
Scan(&role).Error
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return role, nil
|
||||
}
|
||||
|
||||
// TeamRoleRank 角色权重:owner=3 admin=2 member=1。
|
||||
func TeamRoleRank(role string) int {
|
||||
switch role {
|
||||
case "owner":
|
||||
return 3
|
||||
case "admin":
|
||||
return 2
|
||||
case "member":
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// IsTeamAdmin 该用户是否为团队 owner/admin。
|
||||
func IsTeamAdmin(db *gorm.DB, teamID, userID int64) bool {
|
||||
if teamID <= 0 || userID <= 0 {
|
||||
return false
|
||||
}
|
||||
var n int64
|
||||
db.Table("team_members").
|
||||
Where("team_id = ? AND user_id = ? AND role IN ('owner','admin')", teamID, userID).
|
||||
Count(&n)
|
||||
return n > 0
|
||||
}
|
||||
|
||||
// RequireTeamRole 校验最低角色,返回实际角色。
|
||||
func RequireTeamRole(db *gorm.DB, teamID, userID int64, min string) (string, error) {
|
||||
role, err := TeamRole(db, teamID, userID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if TeamRoleRank(role) < TeamRoleRank(min) {
|
||||
return role, Forbidden("TEAM_FORBIDDEN")
|
||||
}
|
||||
return role, nil
|
||||
}
|
||||
125
internal/commonservice/util.go
Normal file
125
internal/commonservice/util.go
Normal file
@@ -0,0 +1,125 @@
|
||||
package commonservice
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"nl-pms-api/internal/config"
|
||||
)
|
||||
|
||||
// AdminUserID 是 code_count 体系的超级管理员账号(users.id=1)。
|
||||
const AdminUserID int64 = 1
|
||||
|
||||
// ParseID 把字符串解析为非负 int64;非法或负数视为 0。
|
||||
func ParseID(s string) int64 {
|
||||
n, _ := strconv.ParseInt(s, 10, 64)
|
||||
if n < 0 {
|
||||
return 0
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// Clip 按字节截断字符串到最多 n 字节。
|
||||
func Clip(s string, n int) string {
|
||||
if len(s) > n {
|
||||
return s[:n]
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// ClipRunes 按 rune 截断到最多 n 个字符。
|
||||
func ClipRunes(s string, n int) string {
|
||||
r := []rune(strings.TrimSpace(s))
|
||||
if len(r) > n {
|
||||
return string(r[:n])
|
||||
}
|
||||
return string(r)
|
||||
}
|
||||
|
||||
// NowRFC 返回当前 UTC 的 RFC3339 时间串(与 view 库惯例一致)。
|
||||
func NowRFC() string {
|
||||
return time.Now().UTC().Format(time.RFC3339)
|
||||
}
|
||||
|
||||
// PublicURL 拼接文件公开访问地址:优先配置的 base_url,否则用 requestHost(可含 scheme://host)。
|
||||
func PublicURL(cfg *config.Config, requestHost, name string) string {
|
||||
base := cfg.BaseURL
|
||||
if base == "" {
|
||||
base = strings.TrimRight(requestHost, "/")
|
||||
if !strings.Contains(base, "://") {
|
||||
base = "http://" + base
|
||||
}
|
||||
}
|
||||
return base + "/files/" + name
|
||||
}
|
||||
|
||||
// NormalizeKind 规范化文件用途:avatar | content,其余为空。
|
||||
func NormalizeKind(k string) string {
|
||||
if k == "avatar" || k == "content" {
|
||||
return k
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// StoredName 生成不可枚举的存储相对路径(日期目录 + 128 位随机 hex)。
|
||||
func StoredName(ext string) string {
|
||||
b := make([]byte, 16)
|
||||
_, _ = rand.Read(b)
|
||||
return time.Now().UTC().Format("2006/01/02") + "/" + hex.EncodeToString(b) + ext
|
||||
}
|
||||
|
||||
// MapStr 从通用 JSON map 取字符串字段。
|
||||
func MapStr(m map[string]any, key string) string {
|
||||
v, ok := m[key]
|
||||
if !ok || v == nil {
|
||||
return ""
|
||||
}
|
||||
switch x := v.(type) {
|
||||
case string:
|
||||
return x
|
||||
case []byte:
|
||||
return string(x)
|
||||
case float64:
|
||||
if x == float64(int64(x)) {
|
||||
return strconv.FormatInt(int64(x), 10)
|
||||
}
|
||||
return strconv.FormatFloat(x, 'f', -1, 64)
|
||||
case bool:
|
||||
if x {
|
||||
return "1"
|
||||
}
|
||||
return "0"
|
||||
default:
|
||||
return fmt.Sprint(x)
|
||||
}
|
||||
}
|
||||
|
||||
// MapInt64 从通用 JSON map 取 int64 字段。
|
||||
func MapInt64(m map[string]any, key string) int64 {
|
||||
v, ok := m[key]
|
||||
if !ok || v == nil {
|
||||
return 0
|
||||
}
|
||||
switch x := v.(type) {
|
||||
case float64:
|
||||
return int64(x)
|
||||
case int64:
|
||||
return x
|
||||
case int:
|
||||
return int64(x)
|
||||
case string:
|
||||
n, _ := strconv.ParseInt(x, 10, 64)
|
||||
return n
|
||||
case bool:
|
||||
if x {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user