* test(relayconvert): add golden snapshot matrix and relaykit boundary guard Phase 0 of the relaykit extraction plan: pin byte-level output of every registered (from,to) request/response/stream conversion route, and forbid kit-bound packages from growing host-only imports. * wip(relayconvert): drop gin.Context from converter signatures; add convmeta draft Phase 1 in progress: relayconvert now takes context.Context; host media resolver adapts gin.Context back at the service boundary. * refactor(relayconvert): decouple converters from RelayInfo, gin, and settings Phase 1 of the relaykit extraction plan: - converters now depend on convmeta.Meta (implemented by RelayInfo) instead of *relaycommon.RelayInfo; ClaudeConvertInfo and the format guesser move to convmeta with aliases left behind - host settings reach converters via a convmeta.Options snapshot built in RelayInfo.ConvOptions; no more model_setting/reasoning global reads inside the conversion layer - effort-suffix helpers move to service/relayconvert/reasoning (old package forwards); chat-to-responses upgrade policy moves to service (host routing logic, not conversion) - golden conversion matrix unchanged * test(relayconvert): tighten boundary — kit packages now free of gin/setting imports * refactor(dto): drop gin and logger dependencies Phase 2 (part 1): dto.Request.IsStream now takes *http.Request instead of *gin.Context (Gemini's impl reads query/path off the std request); dto's three logger calls become common.SysError. Boundary test allowlist is now empty — kit-bound packages import no gin/setting/logger/model. * refactor(kit): extract dependency-free kitutil; dto/types/relayconvert stop importing common Phase 2 of the relaykit extraction plan: - new service/relayconvert/kitutil holds the pure helpers the kit needs (JSON wrappers, pointer/string/uuid/timestamp utils, MaskSensitiveInfo, pluggable LogInfo/LogError hooks, Debug flag) - dto, types, and all relayconvert packages now use kitutil; their only remaining internal deps are dto/types/constant - common keeps every original symbol (MaskSensitiveInfo delegates to kitutil) so host code is untouched; main.go routes kit logging into common.SysLog/SysError and mirrors DebugEnabled - golden conversion matrix unchanged * refactor(kit): move EndpointType/FinishReason to types; OpenRouter dialect via Options Kit packages (dto/types/relayconvert/reasonmap) no longer import constant: - EndpointType and finish-reason values live in types; constant re-exports - the OpenRouter special-case in claude->openai request conversion reads Options.OpenRouterDialect, set by the host from the channel type; InitChannelMeta invalidates the cached snapshot on channel switch * refactor: extract relaykit submodule (dto/types/relayconvert/reasonmap) Phase 3 of the relaykit extraction plan: - new go module github.com/QuantumNous/new-api/relaykit containing dto (minus task family), types, relayconvert (with convmeta/kitutil/reasoning), and reasonmap; host consumes it via require + replace, go.work for dev - task-family dto (task/suno/midjourney/video) stays in the host dto package; dual-consumer host files alias it as taskdto - relaykit builds and tests standalone (GOWORK=off): no host imports, no gin, no DB, no settings - golden conversion matrix unchanged * build(docker): copy relaykit/go.mod before go mod download The local-replace submodule's go.mod must exist inside the build context for the main module graph to resolve. * fix: address relaykit extraction regressions * fix: address relaykit review regressions * docs: document Meta nil receiver contract * fix(relaykit): fail OpenAI→Claude conversion without max_tokens; reject negative default_max_tokens The Claude Messages API requires max_tokens (omitting it is a 400 "Field required"), but with a nil Options.Claude.DefaultMaxTokens hook the converters silently emitted a request the upstream is guaranteed to reject. Both OpenAI Chat and Responses → Claude conversions now return sharedclaude.ErrMissingMaxTokens when no path (client value, default hook, thinking-adapter floor) supplied one. Unreachable in the host, which always configures the hook. Host side, claude.default_max_tokens now rejects negative values at the option API before persisting — they would wrap into huge unsigned values during conversion. Zero stays allowed: the current API treats max_tokens: 0 as cache pre-warming. * fix: make Gemini safety settings read path race-free
516 lines
17 KiB
Go
516 lines
17 KiB
Go
package middleware
|
||
|
||
import (
|
||
"errors"
|
||
"fmt"
|
||
"net"
|
||
"net/http"
|
||
"strings"
|
||
|
||
"github.com/QuantumNous/new-api/common"
|
||
"github.com/QuantumNous/new-api/constant"
|
||
"github.com/QuantumNous/new-api/i18n"
|
||
"github.com/QuantumNous/new-api/logger"
|
||
"github.com/QuantumNous/new-api/model"
|
||
"github.com/QuantumNous/new-api/relaykit/types"
|
||
"github.com/QuantumNous/new-api/service"
|
||
"github.com/QuantumNous/new-api/service/authz"
|
||
"github.com/QuantumNous/new-api/setting/ratio_setting"
|
||
|
||
"github.com/gin-gonic/gin"
|
||
"gorm.io/gorm"
|
||
)
|
||
|
||
const authIdentityContextKey = "auth_identity"
|
||
|
||
type dashboardCredentialKind int
|
||
|
||
const (
|
||
dashboardCredentialUnmatched dashboardCredentialKind = iota
|
||
dashboardCredentialInternal
|
||
dashboardCredentialPAT
|
||
)
|
||
|
||
func validUserInfo(username string, role int) bool {
|
||
// check username is empty
|
||
if strings.TrimSpace(username) == "" {
|
||
return false
|
||
}
|
||
if !common.IsValidateRole(role) {
|
||
return false
|
||
}
|
||
return true
|
||
}
|
||
|
||
func authHelper(c *gin.Context, minRole int) {
|
||
user, identity, useAccessToken, err := authenticateDashboardRequest(c)
|
||
if err != nil {
|
||
writeDashboardAuthError(c, err)
|
||
return
|
||
}
|
||
if user.Status != common.UserStatusEnabled {
|
||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"success": false, "code": "AUTH_USER_DISABLED", "message": common.TranslateMessage(c, i18n.MsgAuthUserBanned)})
|
||
return
|
||
}
|
||
if user.Role < minRole {
|
||
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"success": false, "code": "AUTH_INSUFFICIENT_PRIVILEGE", "message": common.TranslateMessage(c, i18n.MsgAuthInsufficientPrivilege)})
|
||
return
|
||
}
|
||
if !validUserInfo(user.Username, user.Role) {
|
||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"success": false, "code": "AUTH_USER_INVALID", "message": common.TranslateMessage(c, i18n.MsgAuthUserInfoInvalid)})
|
||
return
|
||
}
|
||
setDashboardAuthContext(c, user, identity, useAccessToken)
|
||
|
||
// 管理/root 写操作审计兜底:内聚在鉴权链路里,保证任何经过 AdminAuth/RootAuth
|
||
// 的写接口都会自动留痕(无需在路由上单独挂审计中间件,避免漏挂)。
|
||
// handler 内手动埋点者会设置 ContextKeyAuditLogged,finishAdminAudit 据此跳过。
|
||
var auditWriter *auditResponseWriter
|
||
if minRole >= common.RoleAdminUser {
|
||
auditWriter = beginAdminAudit(c)
|
||
}
|
||
|
||
c.Next()
|
||
|
||
finishAdminAudit(c, auditWriter)
|
||
}
|
||
|
||
func TryUserAuth() func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
user, identity, credentialKind, err := classifyDashboardCredential(c)
|
||
if err != nil {
|
||
writeDashboardAuthError(c, err)
|
||
return
|
||
}
|
||
if credentialKind != dashboardCredentialUnmatched {
|
||
setDashboardAuthContext(c, user, identity, credentialKind == dashboardCredentialPAT)
|
||
}
|
||
c.Next()
|
||
}
|
||
}
|
||
|
||
func UserAuth() func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
authHelper(c, common.RoleCommonUser)
|
||
}
|
||
}
|
||
|
||
func AdminAuth() func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
authHelper(c, common.RoleAdminUser)
|
||
}
|
||
}
|
||
|
||
func RootAuth() func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
authHelper(c, common.RoleRootUser)
|
||
}
|
||
}
|
||
|
||
// GetAuthIdentity returns a dashboard session identity. PAT-authenticated
|
||
// requests intentionally have no SessionID and cannot manage browser sessions.
|
||
func GetAuthIdentity(c *gin.Context) (service.AuthIdentity, bool) {
|
||
value, ok := c.Get(authIdentityContextKey)
|
||
if !ok {
|
||
return service.AuthIdentity{}, false
|
||
}
|
||
identity, ok := value.(service.AuthIdentity)
|
||
return identity, ok
|
||
}
|
||
|
||
// GetSessionAuthIdentity returns only identities backed by a live dashboard
|
||
// session. PAT-authenticated requests intentionally fail this check.
|
||
func GetSessionAuthIdentity(c *gin.Context) (service.AuthIdentity, bool) {
|
||
identity, ok := GetAuthIdentity(c)
|
||
if !ok {
|
||
identity = service.AuthIdentity{
|
||
UserID: c.GetInt("id"),
|
||
SessionID: c.GetString("session_id"),
|
||
UserAuthVersion: c.GetInt64("auth_version"),
|
||
SessionVersion: c.GetInt64("session_version"),
|
||
}
|
||
}
|
||
if identity.UserID <= 0 || identity.SessionID == "" || identity.UserAuthVersion <= 0 || identity.SessionVersion <= 0 {
|
||
return service.AuthIdentity{}, false
|
||
}
|
||
return identity, true
|
||
}
|
||
|
||
func authenticateDashboardRequest(c *gin.Context) (*model.UserBase, service.AuthIdentity, bool, error) {
|
||
user, identity, credentialKind, err := classifyDashboardCredential(c)
|
||
if err != nil {
|
||
return nil, service.AuthIdentity{}, credentialKind == dashboardCredentialPAT, err
|
||
}
|
||
if credentialKind == dashboardCredentialUnmatched {
|
||
return nil, service.AuthIdentity{}, false, service.ErrAuthTokenInvalid
|
||
}
|
||
return user, identity, credentialKind == dashboardCredentialPAT, nil
|
||
}
|
||
|
||
func classifyDashboardCredential(c *gin.Context) (*model.UserBase, service.AuthIdentity, dashboardCredentialKind, error) {
|
||
raw, ok := authorizationToken(c.GetHeader("Authorization"))
|
||
if !ok {
|
||
return nil, service.AuthIdentity{}, dashboardCredentialUnmatched, nil
|
||
}
|
||
identity, internal, err := service.ParseDashboardAccessToken(raw)
|
||
if internal {
|
||
if err != nil {
|
||
return nil, service.AuthIdentity{}, dashboardCredentialInternal, err
|
||
}
|
||
_, user, err := service.ValidateLoginSession(identity)
|
||
if err != nil {
|
||
return nil, service.AuthIdentity{}, dashboardCredentialInternal, err
|
||
}
|
||
return user, identity, dashboardCredentialInternal, nil
|
||
}
|
||
patUser, err := model.ValidateAccessToken(raw)
|
||
if err != nil {
|
||
return nil, service.AuthIdentity{}, dashboardCredentialPAT, err
|
||
}
|
||
if patUser == nil || patUser.Id <= 0 {
|
||
return nil, service.AuthIdentity{}, dashboardCredentialUnmatched, nil
|
||
}
|
||
user, err := model.GetUserCache(patUser.Id)
|
||
if err != nil {
|
||
return nil, service.AuthIdentity{}, dashboardCredentialPAT, err
|
||
}
|
||
return user, service.AuthIdentity{UserID: user.Id, UserAuthVersion: user.AuthVersion}, dashboardCredentialPAT, nil
|
||
}
|
||
|
||
func authorizationToken(header string) (string, bool) {
|
||
header = strings.TrimSpace(header)
|
||
if header == "" {
|
||
return "", false
|
||
}
|
||
parts := strings.Fields(header)
|
||
if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
|
||
header = parts[1]
|
||
} else if len(parts) != 1 {
|
||
return "", false
|
||
}
|
||
return header, header != ""
|
||
}
|
||
|
||
func setDashboardAuthContext(c *gin.Context, user *model.UserBase, identity service.AuthIdentity, useAccessToken bool) {
|
||
c.Header("Auth-Version", "864b7076dbcd0a3c01b5520316720ebf")
|
||
c.Set("username", user.Username)
|
||
c.Set("role", user.Role)
|
||
c.Set("id", user.Id)
|
||
c.Set("group", user.Group)
|
||
c.Set("user_group", user.Group)
|
||
c.Set("use_access_token", useAccessToken)
|
||
c.Set("session_id", identity.SessionID)
|
||
c.Set("auth_version", identity.UserAuthVersion)
|
||
c.Set("session_version", identity.SessionVersion)
|
||
c.Set(authIdentityContextKey, identity)
|
||
user.WriteContext(c)
|
||
}
|
||
|
||
func writeDashboardAuthError(c *gin.Context, err error) {
|
||
if errors.Is(err, service.ErrAuthTokenExpired) {
|
||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"success": false, "code": "AUTH_TOKEN_EXPIRED", "message": common.TranslateMessage(c, i18n.MsgAuthNotLoggedIn)})
|
||
return
|
||
}
|
||
if errors.Is(err, service.ErrLoginSessionRevoked) || errors.Is(err, gorm.ErrRecordNotFound) {
|
||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"success": false, "code": "AUTH_SESSION_REVOKED", "message": common.TranslateMessage(c, i18n.MsgAuthNotLoggedIn)})
|
||
return
|
||
}
|
||
if errors.Is(err, service.ErrAuthTokenInvalid) {
|
||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"success": false, "code": "AUTH_UNAUTHORIZED", "message": common.TranslateMessage(c, i18n.MsgAuthAccessTokenInvalid)})
|
||
return
|
||
}
|
||
common.SysLog("dashboard authentication error: " + err.Error())
|
||
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"success": false, "code": "AUTH_INTERNAL_ERROR", "message": common.TranslateMessage(c, i18n.MsgDatabaseError)})
|
||
}
|
||
|
||
func RequirePermission(permission authz.Permission) func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
role := c.GetInt("role")
|
||
userID := c.GetInt("id")
|
||
if authz.Can(userID, role, permission) {
|
||
c.Next()
|
||
return
|
||
}
|
||
c.JSON(http.StatusForbidden, gin.H{
|
||
"success": false,
|
||
"message": common.TranslateMessage(c, i18n.MsgAuthInsufficientPrivilege),
|
||
})
|
||
c.Abort()
|
||
}
|
||
}
|
||
|
||
func WssAuth(c *gin.Context) {
|
||
|
||
}
|
||
|
||
// TokenOrUserAuth allows either session-based user auth or API token auth.
|
||
// Used for endpoints that need to be accessible from both the dashboard and API clients.
|
||
func TokenOrUserAuth() func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
raw, ok := authorizationToken(c.GetHeader("Authorization"))
|
||
if ok {
|
||
identity, internal, err := service.ParseDashboardAccessToken(raw)
|
||
if !internal {
|
||
TokenAuth()(c)
|
||
return
|
||
}
|
||
if err != nil {
|
||
writeDashboardAuthError(c, err)
|
||
return
|
||
}
|
||
_, user, err := service.ValidateLoginSession(identity)
|
||
if err != nil {
|
||
writeDashboardAuthError(c, err)
|
||
return
|
||
}
|
||
setDashboardAuthContext(c, user, identity, false)
|
||
c.Next()
|
||
return
|
||
}
|
||
// Opaque credentials are relay API keys here, never dashboard PATs.
|
||
TokenAuth()(c)
|
||
}
|
||
}
|
||
|
||
// TokenAuthReadOnly 宽松版本的令牌认证中间件,用于只读查询接口。
|
||
// 只验证令牌 key 是否存在,不检查令牌状态、过期时间和额度。
|
||
// 即使令牌已过期、已耗尽或已禁用,也允许访问。
|
||
// 仍然检查用户是否被封禁。
|
||
func TokenAuthReadOnly() func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
key := c.Request.Header.Get("Authorization")
|
||
if key == "" {
|
||
c.JSON(http.StatusUnauthorized, gin.H{
|
||
"success": false,
|
||
"message": common.TranslateMessage(c, i18n.MsgTokenNotProvided),
|
||
})
|
||
c.Abort()
|
||
return
|
||
}
|
||
if strings.HasPrefix(key, "Bearer ") || strings.HasPrefix(key, "bearer ") {
|
||
key = strings.TrimSpace(key[7:])
|
||
}
|
||
key = strings.TrimPrefix(key, "sk-")
|
||
parts := strings.Split(key, "-")
|
||
key = parts[0]
|
||
|
||
token, err := model.GetTokenByKey(key, false)
|
||
if err != nil {
|
||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||
c.JSON(http.StatusUnauthorized, gin.H{
|
||
"success": false,
|
||
"message": common.TranslateMessage(c, i18n.MsgTokenInvalid),
|
||
})
|
||
} else {
|
||
common.SysLog("TokenAuthReadOnly GetTokenByKey database error: " + err.Error())
|
||
c.JSON(http.StatusInternalServerError, gin.H{
|
||
"success": false,
|
||
"message": common.TranslateMessage(c, i18n.MsgDatabaseError),
|
||
})
|
||
}
|
||
c.Abort()
|
||
return
|
||
}
|
||
|
||
// TokenAuthReadOnly must keep allowing other token states to query read-only
|
||
// data, such as token usage logs; only explicitly disabled tokens are denied.
|
||
if token.Status == common.TokenStatusDisabled {
|
||
c.JSON(http.StatusUnauthorized, gin.H{
|
||
"success": false,
|
||
"message": common.TranslateMessage(c, i18n.MsgTokenStatusUnavailable),
|
||
})
|
||
c.Abort()
|
||
return
|
||
}
|
||
|
||
userCache, err := model.GetUserCache(token.UserId)
|
||
if err != nil {
|
||
common.SysLog(fmt.Sprintf("TokenAuthReadOnly GetUserCache error for user %d: %v", token.UserId, err))
|
||
c.JSON(http.StatusInternalServerError, gin.H{
|
||
"success": false,
|
||
"message": common.TranslateMessage(c, i18n.MsgDatabaseError),
|
||
})
|
||
c.Abort()
|
||
return
|
||
}
|
||
if userCache.Status != common.UserStatusEnabled {
|
||
c.JSON(http.StatusForbidden, gin.H{
|
||
"success": false,
|
||
"message": common.TranslateMessage(c, i18n.MsgAuthUserBanned),
|
||
})
|
||
c.Abort()
|
||
return
|
||
}
|
||
|
||
c.Set("id", token.UserId)
|
||
c.Set("token_id", token.Id)
|
||
c.Set("token_key", token.Key)
|
||
c.Next()
|
||
}
|
||
}
|
||
|
||
func TokenAuth() func(c *gin.Context) {
|
||
return func(c *gin.Context) {
|
||
// 先检测是否为ws
|
||
if c.Request.Header.Get("Sec-WebSocket-Protocol") != "" {
|
||
// Sec-WebSocket-Protocol: realtime, openai-insecure-api-key.sk-xxx, openai-beta.realtime-v1
|
||
// read sk from Sec-WebSocket-Protocol
|
||
key := c.Request.Header.Get("Sec-WebSocket-Protocol")
|
||
parts := strings.Split(key, ",")
|
||
for _, part := range parts {
|
||
part = strings.TrimSpace(part)
|
||
if strings.HasPrefix(part, "openai-insecure-api-key") {
|
||
key = strings.TrimPrefix(part, "openai-insecure-api-key.")
|
||
break
|
||
}
|
||
}
|
||
c.Request.Header.Set("Authorization", "Bearer "+key)
|
||
}
|
||
// 检查path包含/v1/messages 或 /v1/models
|
||
if strings.Contains(c.Request.URL.Path, "/v1/messages") || strings.Contains(c.Request.URL.Path, "/v1/models") {
|
||
anthropicKey := c.Request.Header.Get("x-api-key")
|
||
if anthropicKey != "" {
|
||
c.Request.Header.Set("Authorization", "Bearer "+anthropicKey)
|
||
}
|
||
}
|
||
// gemini api 从query中获取key
|
||
if strings.HasPrefix(c.Request.URL.Path, "/v1beta/models") ||
|
||
strings.HasPrefix(c.Request.URL.Path, "/v1beta/openai/models") ||
|
||
strings.HasPrefix(c.Request.URL.Path, "/v1/models/") {
|
||
skKey := c.Query("key")
|
||
if skKey != "" {
|
||
c.Request.Header.Set("Authorization", "Bearer "+skKey)
|
||
}
|
||
// 从x-goog-api-key header中获取key
|
||
xGoogKey := c.Request.Header.Get("x-goog-api-key")
|
||
if xGoogKey != "" {
|
||
c.Request.Header.Set("Authorization", "Bearer "+xGoogKey)
|
||
}
|
||
}
|
||
key := c.Request.Header.Get("Authorization")
|
||
parts := make([]string, 0)
|
||
if strings.HasPrefix(key, "Bearer ") || strings.HasPrefix(key, "bearer ") {
|
||
key = strings.TrimSpace(key[7:])
|
||
}
|
||
if key == "" || key == "midjourney-proxy" {
|
||
key = c.Request.Header.Get("mj-api-secret")
|
||
if strings.HasPrefix(key, "Bearer ") || strings.HasPrefix(key, "bearer ") {
|
||
key = strings.TrimSpace(key[7:])
|
||
}
|
||
key = strings.TrimPrefix(key, "sk-")
|
||
parts = strings.Split(key, "-")
|
||
key = parts[0]
|
||
} else {
|
||
key = strings.TrimPrefix(key, "sk-")
|
||
parts = strings.Split(key, "-")
|
||
key = parts[0]
|
||
}
|
||
token, err := model.ValidateUserToken(key)
|
||
if token != nil {
|
||
id := c.GetInt("id")
|
||
if id == 0 {
|
||
c.Set("id", token.UserId)
|
||
}
|
||
}
|
||
if err != nil {
|
||
if errors.Is(err, model.ErrDatabase) {
|
||
common.SysLog("TokenAuth ValidateUserToken database error: " + err.Error())
|
||
abortWithOpenAiMessage(c, http.StatusInternalServerError,
|
||
common.TranslateMessage(c, i18n.MsgDatabaseError))
|
||
} else {
|
||
abortWithOpenAiMessage(c, http.StatusUnauthorized,
|
||
common.TranslateMessage(c, i18n.MsgTokenInvalid))
|
||
}
|
||
return
|
||
}
|
||
|
||
allowIps := token.GetIpLimits()
|
||
if len(allowIps) > 0 {
|
||
clientIp := c.ClientIP()
|
||
logger.LogDebug(c, "Token has IP restrictions, checking client IP %s", clientIp)
|
||
ip := net.ParseIP(clientIp)
|
||
if ip == nil {
|
||
abortWithOpenAiMessage(c, http.StatusForbidden, "无法解析客户端 IP 地址")
|
||
return
|
||
}
|
||
if common.IsIpInCIDRList(ip, allowIps) == false {
|
||
abortWithOpenAiMessage(c, http.StatusForbidden, "您的 IP 不在令牌允许访问的列表中", types.ErrorCodeAccessDenied)
|
||
return
|
||
}
|
||
logger.LogDebug(c, "Client IP %s passed the token IP restrictions check", clientIp)
|
||
}
|
||
|
||
userCache, err := model.GetUserCache(token.UserId)
|
||
if err != nil {
|
||
common.SysLog(fmt.Sprintf("TokenAuth GetUserCache error for user %d: %v", token.UserId, err))
|
||
abortWithOpenAiMessage(c, http.StatusInternalServerError,
|
||
common.TranslateMessage(c, i18n.MsgDatabaseError))
|
||
return
|
||
}
|
||
userEnabled := userCache.Status == common.UserStatusEnabled
|
||
if !userEnabled {
|
||
abortWithOpenAiMessage(c, http.StatusForbidden, common.TranslateMessage(c, i18n.MsgAuthUserBanned))
|
||
return
|
||
}
|
||
|
||
userCache.WriteContext(c)
|
||
|
||
userGroup := userCache.Group
|
||
tokenGroup := token.Group
|
||
if tokenGroup != "" {
|
||
// check common.UserUsableGroups[userGroup]
|
||
if _, ok := service.GetUserUsableGroups(userGroup)[tokenGroup]; !ok {
|
||
abortWithOpenAiMessage(c, http.StatusForbidden, fmt.Sprintf("无权访问 %s 分组", tokenGroup))
|
||
return
|
||
}
|
||
// check group in common.GroupRatio
|
||
if !ratio_setting.ContainsGroupRatio(tokenGroup) {
|
||
if tokenGroup != "auto" {
|
||
abortWithOpenAiMessage(c, http.StatusForbidden, fmt.Sprintf("分组 %s 已被弃用", tokenGroup))
|
||
return
|
||
}
|
||
}
|
||
userGroup = tokenGroup
|
||
}
|
||
common.SetContextKey(c, constant.ContextKeyUsingGroup, userGroup)
|
||
|
||
err = SetupContextForToken(c, token, parts...)
|
||
if err != nil {
|
||
return
|
||
}
|
||
c.Next()
|
||
}
|
||
}
|
||
|
||
func SetupContextForToken(c *gin.Context, token *model.Token, parts ...string) error {
|
||
if token == nil {
|
||
return fmt.Errorf("token is nil")
|
||
}
|
||
c.Set("id", token.UserId)
|
||
c.Set("token_id", token.Id)
|
||
c.Set("token_key", token.Key)
|
||
c.Set("token_name", token.Name)
|
||
c.Set("token_unlimited_quota", token.UnlimitedQuota)
|
||
if !token.UnlimitedQuota {
|
||
c.Set("token_quota", token.RemainQuota)
|
||
}
|
||
if token.ModelLimitsEnabled {
|
||
c.Set("token_model_limit_enabled", true)
|
||
c.Set("token_model_limit", token.GetModelLimitsMap())
|
||
} else {
|
||
c.Set("token_model_limit_enabled", false)
|
||
}
|
||
common.SetContextKey(c, constant.ContextKeyTokenGroup, token.Group)
|
||
common.SetContextKey(c, constant.ContextKeyTokenCrossGroupRetry, token.CrossGroupRetry)
|
||
if len(parts) > 1 {
|
||
if model.IsAdmin(token.UserId) {
|
||
c.Set("specific_channel_id", parts[1])
|
||
} else {
|
||
c.Header("specific_channel_version", "701e3ae1dc3f7975556d354e0675168d004891c8")
|
||
abortWithOpenAiMessage(c, http.StatusForbidden, "普通用户不支持指定渠道")
|
||
return fmt.Errorf("普通用户不支持指定渠道")
|
||
}
|
||
}
|
||
return nil
|
||
}
|