Compare commits
10 Commits
5b0b69f17e
...
fd3e95ac29
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fd3e95ac29 | ||
|
|
0f90d2b725 | ||
|
|
e8642f4bba | ||
|
|
67c173cbcb | ||
|
|
3f24ba87ee | ||
|
|
13da2edb0a | ||
|
|
4eb7373033 | ||
|
|
0427cd067a | ||
|
|
1612c30de6 | ||
|
|
769a700b5c |
@ -56,6 +56,12 @@ const (
|
|||||||
|
|
||||||
ContextKeySystemPromptOverride ContextKey = "system_prompt_override"
|
ContextKeySystemPromptOverride ContextKey = "system_prompt_override"
|
||||||
|
|
||||||
|
ContextKeyMuseUserID ContextKey = "muse_user_id"
|
||||||
|
ContextKeyMuseWorkID ContextKey = "muse_work_id"
|
||||||
|
ContextKeyMuseRequestID ContextKey = "muse_request_id"
|
||||||
|
ContextKeyMuseScene ContextKey = "muse_scene"
|
||||||
|
ContextKeyMuseTraceID ContextKey = "muse_trace_id"
|
||||||
|
|
||||||
// ContextKeyFileSourcesToCleanup stores file sources that need cleanup when request ends
|
// ContextKeyFileSourcesToCleanup stores file sources that need cleanup when request ends
|
||||||
ContextKeyFileSourcesToCleanup ContextKey = "file_sources_to_cleanup"
|
ContextKeyFileSourcesToCleanup ContextKey = "file_sources_to_cleanup"
|
||||||
|
|
||||||
|
|||||||
@ -20,6 +20,7 @@ func GetAllLogs(c *gin.Context) {
|
|||||||
modelName := c.Query("model_name")
|
modelName := c.Query("model_name")
|
||||||
channel, _ := strconv.Atoi(c.Query("channel"))
|
channel, _ := strconv.Atoi(c.Query("channel"))
|
||||||
group := c.Query("group")
|
group := c.Query("group")
|
||||||
|
// request_id filters the internal relay request id; muse_request_id is carried in log.other for drill-down.
|
||||||
requestId := c.Query("request_id")
|
requestId := c.Query("request_id")
|
||||||
logs, total, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId)
|
logs, total, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@ -41,6 +42,7 @@ func GetUserLogs(c *gin.Context) {
|
|||||||
tokenName := c.Query("token_name")
|
tokenName := c.Query("token_name")
|
||||||
modelName := c.Query("model_name")
|
modelName := c.Query("model_name")
|
||||||
group := c.Query("group")
|
group := c.Query("group")
|
||||||
|
// request_id filters the internal relay request id; muse_request_id is carried in log.other for drill-down.
|
||||||
requestId := c.Query("request_id")
|
requestId := c.Query("request_id")
|
||||||
logs, total, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId)
|
logs, total, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
67
controller/muse_internal.go
Normal file
67
controller/muse_internal.go
Normal file
@ -0,0 +1,67 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"strconv"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
museDTO "github.com/QuantumNous/new-api/dto"
|
||||||
|
"github.com/QuantumNous/new-api/service"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
func SyncMuseUser(c *gin.Context) {
|
||||||
|
var req museDTO.SyncMuseUserRequest
|
||||||
|
if err := common.DecodeJson(c.Request.Body, &req); err != nil {
|
||||||
|
common.ApiError(c, errors.New("invalid request body"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
response, err := service.SyncMuseUser(req)
|
||||||
|
if err != nil {
|
||||||
|
common.ApiError(c, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
common.ApiSuccess(c, response)
|
||||||
|
}
|
||||||
|
|
||||||
|
func DisableMuseUser(c *gin.Context) {
|
||||||
|
userID, err := strconv.Atoi(c.Param("id"))
|
||||||
|
if err != nil {
|
||||||
|
common.ApiError(c, errors.New("invalid user id"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
response, err := service.DisableMuseUserByID(userID)
|
||||||
|
if err != nil {
|
||||||
|
common.ApiError(c, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
common.ApiSuccess(c, response)
|
||||||
|
}
|
||||||
|
|
||||||
|
func RevokeMuseToken(c *gin.Context) {
|
||||||
|
tokenID, err := strconv.Atoi(c.Param("id"))
|
||||||
|
if err != nil {
|
||||||
|
common.ApiError(c, errors.New("invalid token id"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
response, err := service.RevokeMuseToken(tokenID)
|
||||||
|
if err != nil {
|
||||||
|
common.ApiError(c, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
common.ApiSuccess(c, response)
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetMuseUserStatus(c *gin.Context) {
|
||||||
|
userID, err := strconv.Atoi(c.Param("id"))
|
||||||
|
if err != nil {
|
||||||
|
common.ApiError(c, errors.New("invalid user id"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
response, err := service.GetMuseUserStatus(userID)
|
||||||
|
if err != nil {
|
||||||
|
common.ApiError(c, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
common.ApiSuccess(c, response)
|
||||||
|
}
|
||||||
151
controller/muse_internal_test.go
Normal file
151
controller/muse_internal_test.go
Normal file
@ -0,0 +1,151 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
museDTO "github.com/QuantumNous/new-api/dto"
|
||||||
|
"github.com/QuantumNous/new-api/middleware"
|
||||||
|
"github.com/QuantumNous/new-api/model"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/glebarez/sqlite"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
type museInternalAPIResponse struct {
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
Data museDTO.SyncMuseUserResponse `json:"data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupMuseInternalControllerTestDB(t *testing.T) *gorm.DB {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
common.UsingSQLite = true
|
||||||
|
common.UsingMySQL = false
|
||||||
|
common.UsingPostgreSQL = false
|
||||||
|
common.RedisEnabled = false
|
||||||
|
common.GlobalApiRateLimitEnable = false
|
||||||
|
|
||||||
|
t.Setenv("MUSE_INTERNAL_SECRET", "test-secret")
|
||||||
|
|
||||||
|
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_"))
|
||||||
|
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
model.DB = db
|
||||||
|
model.LOG_DB = db
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.User{}, &model.Token{}))
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
if err == nil {
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeMuseInternalResponse(t *testing.T, recorder *httptest.ResponseRecorder) museInternalAPIResponse {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
var response museInternalAPIResponse
|
||||||
|
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
|
||||||
|
return response
|
||||||
|
}
|
||||||
|
|
||||||
|
func newMuseInternalRouter() *gin.Engine {
|
||||||
|
engine := gin.New()
|
||||||
|
internal := engine.Group("/api/internal/muse")
|
||||||
|
internal.Use(middleware.MuseInternalAuth())
|
||||||
|
{
|
||||||
|
internal.POST("/users/sync", SyncMuseUser)
|
||||||
|
internal.POST("/users/:id/disable", DisableMuseUser)
|
||||||
|
internal.POST("/tokens/:id/revoke", RevokeMuseToken)
|
||||||
|
internal.GET("/users/:id/status", GetMuseUserStatus)
|
||||||
|
}
|
||||||
|
return engine
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMuseInternalSyncUserCreatesMirrorUserAndToken(t *testing.T) {
|
||||||
|
setupMuseInternalControllerTestDB(t)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/internal/muse/users/sync", strings.NewReader(`{"muse_user_id":"u-1","username":"alice","display_name":"Alice","status":"active"}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("X-Muse-Service-Secret", "test-secret")
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
|
||||||
|
newMuseInternalRouter().ServeHTTP(rr, req)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, rr.Code)
|
||||||
|
response := decodeMuseInternalResponse(t, rr)
|
||||||
|
require.True(t, response.Success)
|
||||||
|
require.Equal(t, "active", response.Data.Status)
|
||||||
|
require.NotZero(t, response.Data.NewAPIUserID)
|
||||||
|
require.NotZero(t, response.Data.TokenID)
|
||||||
|
require.NotEmpty(t, response.Data.TokenKey)
|
||||||
|
|
||||||
|
user, err := model.GetUserById(response.Data.NewAPIUserID, false)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, common.UserStatusEnabled, user.Status)
|
||||||
|
require.Equal(t, "alice", user.Username)
|
||||||
|
require.Equal(t, "muse_user_id:u-1", user.Remark)
|
||||||
|
|
||||||
|
token, err := model.GetTokenById(response.Data.TokenID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, user.Id, token.UserId)
|
||||||
|
require.Equal(t, common.TokenStatusEnabled, token.Status)
|
||||||
|
require.True(t, token.UnlimitedQuota)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMuseInternalSyncUserRejectsInvalidSecret(t *testing.T) {
|
||||||
|
setupMuseInternalControllerTestDB(t)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/internal/muse/users/sync", strings.NewReader(`{"muse_user_id":"u-1","username":"alice","display_name":"Alice","status":"active"}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("X-Muse-Service-Secret", "wrong-secret")
|
||||||
|
rr := httptest.NewRecorder()
|
||||||
|
|
||||||
|
newMuseInternalRouter().ServeHTTP(rr, req)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusUnauthorized, rr.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMuseInternalDisableUserDisablesUserAndToken(t *testing.T) {
|
||||||
|
setupMuseInternalControllerTestDB(t)
|
||||||
|
|
||||||
|
syncReq := httptest.NewRequest(http.MethodPost, "/api/internal/muse/users/sync", strings.NewReader(`{"muse_user_id":"u-2","username":"bob","display_name":"Bob","status":"active"}`))
|
||||||
|
syncReq.Header.Set("Content-Type", "application/json")
|
||||||
|
syncReq.Header.Set("X-Muse-Service-Secret", "test-secret")
|
||||||
|
syncRR := httptest.NewRecorder()
|
||||||
|
engine := newMuseInternalRouter()
|
||||||
|
engine.ServeHTTP(syncRR, syncReq)
|
||||||
|
|
||||||
|
syncResponse := decodeMuseInternalResponse(t, syncRR)
|
||||||
|
require.True(t, syncResponse.Success)
|
||||||
|
|
||||||
|
disableReq := httptest.NewRequest(http.MethodPost, fmt.Sprintf("/api/internal/muse/users/%d/disable", syncResponse.Data.NewAPIUserID), nil)
|
||||||
|
disableReq.Header.Set("X-Muse-Service-Secret", "test-secret")
|
||||||
|
disableRR := httptest.NewRecorder()
|
||||||
|
engine.ServeHTTP(disableRR, disableReq)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, disableRR.Code)
|
||||||
|
disableResponse := decodeMuseInternalResponse(t, disableRR)
|
||||||
|
require.True(t, disableResponse.Success)
|
||||||
|
require.Equal(t, "disabled", disableResponse.Data.Status)
|
||||||
|
require.Equal(t, syncResponse.Data.NewAPIUserID, disableResponse.Data.NewAPIUserID)
|
||||||
|
|
||||||
|
user, err := model.GetUserById(syncResponse.Data.NewAPIUserID, false)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, common.UserStatusDisabled, user.Status)
|
||||||
|
|
||||||
|
token, err := model.GetTokenById(syncResponse.Data.TokenID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, common.TokenStatusDisabled, token.Status)
|
||||||
|
}
|
||||||
15
dto/muse_internal.go
Normal file
15
dto/muse_internal.go
Normal file
@ -0,0 +1,15 @@
|
|||||||
|
package dto
|
||||||
|
|
||||||
|
type SyncMuseUserRequest struct {
|
||||||
|
MuseUserID string `json:"muse_user_id"`
|
||||||
|
Username string `json:"username"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type SyncMuseUserResponse struct {
|
||||||
|
NewAPIUserID int `json:"newapi_user_id"`
|
||||||
|
TokenID int `json:"token_id"`
|
||||||
|
TokenKey string `json:"token_key,omitempty"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
}
|
||||||
@ -23,8 +23,15 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type ModelRequest struct {
|
type ModelRequest struct {
|
||||||
Model string `json:"model"`
|
Model string `json:"model"`
|
||||||
Group string `json:"group,omitempty"`
|
Group string `json:"group,omitempty"`
|
||||||
|
Metadata *MuseRoutingPolicyInput `json:"metadata,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type MuseRoutingPolicyInput struct {
|
||||||
|
Scene string `json:"scene,omitempty"`
|
||||||
|
RoutingPolicy string `json:"routing_policy,omitempty"`
|
||||||
|
DegradePolicy string `json:"degrade_policy,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func Distribute() func(c *gin.Context) {
|
func Distribute() func(c *gin.Context) {
|
||||||
@ -36,6 +43,12 @@ func Distribute() func(c *gin.Context) {
|
|||||||
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidRequest, map[string]any{"Error": err.Error()}))
|
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidRequest, map[string]any{"Error": err.Error()}))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
policy := service.ResolveMuseRoutingPolicy(
|
||||||
|
firstNonBlank(common.GetContextKeyString(c, constant.ContextKeyMuseScene), modelRequest.MetadataValue("scene")),
|
||||||
|
modelRequest.MetadataValue("routing_policy"),
|
||||||
|
modelRequest.MetadataValue("degrade_policy"),
|
||||||
|
)
|
||||||
|
service.SetMuseRoutingPolicy(c, policy)
|
||||||
if ok {
|
if ok {
|
||||||
id, err := strconv.Atoi(channelId.(string))
|
id, err := strconv.Atoi(channelId.(string))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@ -98,6 +111,7 @@ func Distribute() func(c *gin.Context) {
|
|||||||
common.SetContextKey(c, constant.ContextKeyUsingGroup, usingGroup)
|
common.SetContextKey(c, constant.ContextKeyUsingGroup, usingGroup)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
usingGroup = service.ApplyMuseRoutingPolicy(c, usingGroup)
|
||||||
|
|
||||||
if preferredChannelID, found := service.GetPreferredChannelByAffinity(c, modelRequest.Model, usingGroup); found {
|
if preferredChannelID, found := service.GetPreferredChannelByAffinity(c, modelRequest.Model, usingGroup); found {
|
||||||
preferred, err := model.CacheGetChannel(preferredChannelID)
|
preferred, err := model.CacheGetChannel(preferredChannelID)
|
||||||
@ -109,7 +123,7 @@ func Distribute() func(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
} else if usingGroup == "auto" {
|
} else if usingGroup == "auto" {
|
||||||
userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup)
|
userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup)
|
||||||
autoGroups := service.GetUserAutoGroup(userGroup)
|
autoGroups := service.ResolveMusePreferredGroups(c, userGroup)
|
||||||
for _, g := range autoGroups {
|
for _, g := range autoGroups {
|
||||||
if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, preferred.Id) {
|
if model.IsChannelEnabledForGroupModel(g, modelRequest.Model, preferred.Id) {
|
||||||
selectGroup = g
|
selectGroup = g
|
||||||
@ -128,12 +142,31 @@ func Distribute() func(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if channel == nil {
|
if channel == nil {
|
||||||
channel, selectGroup, err = service.CacheGetRandomSatisfiedChannel(&service.RetryParam{
|
retryParam := &service.RetryParam{
|
||||||
Ctx: c,
|
Ctx: c,
|
||||||
ModelName: modelRequest.Model,
|
ModelName: modelRequest.Model,
|
||||||
TokenGroup: usingGroup,
|
TokenGroup: usingGroup,
|
||||||
Retry: common.GetPointer(0),
|
Retry: common.GetPointer(0),
|
||||||
})
|
}
|
||||||
|
channel, selectGroup, err = service.CacheGetRandomSatisfiedChannel(retryParam)
|
||||||
|
if channel == nil && err == nil {
|
||||||
|
candidates := service.MuseFallbackModelCandidates(modelRequest.Model, policy)
|
||||||
|
for _, candidateModel := range candidates[1:] {
|
||||||
|
candidateParam := &service.RetryParam{
|
||||||
|
Ctx: c,
|
||||||
|
ModelName: candidateModel,
|
||||||
|
TokenGroup: usingGroup,
|
||||||
|
Retry: common.GetPointer(0),
|
||||||
|
}
|
||||||
|
channel, selectGroup, err = service.CacheGetRandomSatisfiedChannel(candidateParam)
|
||||||
|
if channel != nil || err != nil {
|
||||||
|
if channel != nil {
|
||||||
|
modelRequest.Model = candidateModel
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
showGroup := usingGroup
|
showGroup := usingGroup
|
||||||
if usingGroup == "auto" {
|
if usingGroup == "auto" {
|
||||||
@ -164,6 +197,31 @@ func Distribute() func(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *ModelRequest) MetadataValue(name string) string {
|
||||||
|
if r == nil || r.Metadata == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
switch name {
|
||||||
|
case "scene":
|
||||||
|
return r.Metadata.Scene
|
||||||
|
case "routing_policy":
|
||||||
|
return r.Metadata.RoutingPolicy
|
||||||
|
case "degrade_policy":
|
||||||
|
return r.Metadata.DegradePolicy
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func firstNonBlank(values ...string) string {
|
||||||
|
for _, value := range values {
|
||||||
|
if strings.TrimSpace(value) != "" {
|
||||||
|
return strings.TrimSpace(value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
// getModelFromRequest 从请求中读取模型信息
|
// getModelFromRequest 从请求中读取模型信息
|
||||||
// 根据 Content-Type 自动处理:
|
// 根据 Content-Type 自动处理:
|
||||||
// - application/json
|
// - application/json
|
||||||
|
|||||||
25
middleware/muse_internal_auth.go
Normal file
25
middleware/muse_internal_auth.go
Normal file
@ -0,0 +1,25 @@
|
|||||||
|
package middleware
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/subtle"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
func MuseInternalAuth() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
secret := c.GetHeader("X-Muse-Service-Secret")
|
||||||
|
expected := common.GetEnvOrDefaultString("MUSE_INTERNAL_SECRET", "")
|
||||||
|
if expected == "" || subtle.ConstantTimeCompare([]byte(secret), []byte(expected)) != 1 {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{
|
||||||
|
"success": false,
|
||||||
|
"message": "invalid muse internal secret",
|
||||||
|
})
|
||||||
|
c.Abort()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
113
middleware/muse_request_context.go
Normal file
113
middleware/muse_request_context.go
Normal file
@ -0,0 +1,113 @@
|
|||||||
|
package middleware
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/QuantumNous/new-api/constant"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
museHeaderUserID = "X-Muse-User-Id"
|
||||||
|
museHeaderWorkID = "X-Muse-Work-Id"
|
||||||
|
museHeaderRequest = "X-Request-Id"
|
||||||
|
museHeaderTraceID = "X-Trace-Id"
|
||||||
|
museHeaderScene = "X-Muse-Scene"
|
||||||
|
jsonContentType = "application/json"
|
||||||
|
)
|
||||||
|
|
||||||
|
type museMetadataEnvelope struct {
|
||||||
|
Metadata map[string]any `json:"metadata"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func MuseRequestContext() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
metadata := extractMuseMetadata(c)
|
||||||
|
|
||||||
|
museUserID := museFirstNonBlank(
|
||||||
|
c.GetHeader(museHeaderUserID),
|
||||||
|
stringifyMetadata(metadata["muse_user_id"]),
|
||||||
|
)
|
||||||
|
museWorkID := museFirstNonBlank(
|
||||||
|
c.GetHeader(museHeaderWorkID),
|
||||||
|
stringifyMetadata(metadata["work_id"]),
|
||||||
|
stringifyMetadata(metadata["muse_work_id"]),
|
||||||
|
)
|
||||||
|
museScene := museFirstNonBlank(
|
||||||
|
c.GetHeader(museHeaderScene),
|
||||||
|
stringifyMetadata(metadata["scene"]),
|
||||||
|
)
|
||||||
|
requestID := museFirstNonBlank(
|
||||||
|
c.GetHeader(museHeaderRequest),
|
||||||
|
stringifyMetadata(metadata["request_id"]),
|
||||||
|
)
|
||||||
|
traceID := museFirstNonBlank(
|
||||||
|
c.GetHeader(museHeaderTraceID),
|
||||||
|
stringifyMetadata(metadata["trace_id"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
if museUserID != "" {
|
||||||
|
common.SetContextKey(c, constant.ContextKeyMuseUserID, museUserID)
|
||||||
|
}
|
||||||
|
if museWorkID != "" {
|
||||||
|
common.SetContextKey(c, constant.ContextKeyMuseWorkID, museWorkID)
|
||||||
|
}
|
||||||
|
if requestID != "" {
|
||||||
|
common.SetContextKey(c, constant.ContextKeyMuseRequestID, requestID)
|
||||||
|
}
|
||||||
|
if museScene != "" {
|
||||||
|
common.SetContextKey(c, constant.ContextKeyMuseScene, museScene)
|
||||||
|
}
|
||||||
|
if traceID != "" {
|
||||||
|
common.SetContextKey(c, constant.ContextKeyMuseTraceID, traceID)
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractMuseMetadata(c *gin.Context) map[string]any {
|
||||||
|
if c == nil || c.Request == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if !strings.Contains(strings.ToLower(c.GetHeader("Content-Type")), jsonContentType) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if !needsMuseMetadataRead(c) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var envelope museMetadataEnvelope
|
||||||
|
if err := common.UnmarshalBodyReusable(c, &envelope); err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return envelope.Metadata
|
||||||
|
}
|
||||||
|
|
||||||
|
func needsMuseMetadataRead(c *gin.Context) bool {
|
||||||
|
return c.GetHeader(museHeaderUserID) == "" ||
|
||||||
|
c.GetHeader(museHeaderWorkID) == "" ||
|
||||||
|
c.GetHeader(museHeaderRequest) == "" ||
|
||||||
|
c.GetHeader(museHeaderTraceID) == "" ||
|
||||||
|
c.GetHeader(museHeaderScene) == ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func museFirstNonBlank(values ...string) string {
|
||||||
|
for _, value := range values {
|
||||||
|
if strings.TrimSpace(value) != "" {
|
||||||
|
return strings.TrimSpace(value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func stringifyMetadata(value any) string {
|
||||||
|
switch v := value.(type) {
|
||||||
|
case nil:
|
||||||
|
return ""
|
||||||
|
case string:
|
||||||
|
return strings.TrimSpace(v)
|
||||||
|
default:
|
||||||
|
return strings.TrimSpace(fmt.Sprintf("%v", v))
|
||||||
|
}
|
||||||
|
}
|
||||||
176
middleware/muse_request_context_test.go
Normal file
176
middleware/muse_request_context_test.go
Normal file
@ -0,0 +1,176 @@
|
|||||||
|
package middleware
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/QuantumNous/new-api/constant"
|
||||||
|
"github.com/QuantumNous/new-api/model"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/glebarez/sqlite"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
type museRequestContextResponse struct {
|
||||||
|
MuseUserID string `json:"muse_user_id"`
|
||||||
|
MuseWorkID string `json:"muse_work_id"`
|
||||||
|
MuseScene string `json:"muse_scene"`
|
||||||
|
MuseRequestID string `json:"muse_request_id"`
|
||||||
|
TraceID string `json:"trace_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupMuseRequestContextTestDB(t *testing.T) *gorm.DB {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
common.UsingSQLite = true
|
||||||
|
common.UsingMySQL = false
|
||||||
|
common.UsingPostgreSQL = false
|
||||||
|
common.RedisEnabled = false
|
||||||
|
|
||||||
|
db, err := gorm.Open(sqlite.Open("file:muse_request_context_test?mode=memory&cache=shared"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
model.DB = db
|
||||||
|
model.LOG_DB = db
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.User{}, &model.Log{}))
|
||||||
|
require.NoError(t, db.Create(&model.User{
|
||||||
|
Id: 1,
|
||||||
|
Username: "alice",
|
||||||
|
Password: "hashed",
|
||||||
|
Status: common.UserStatusEnabled,
|
||||||
|
}).Error)
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
if err == nil {
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMuseRequestContextExtractsHeaders(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPost,
|
||||||
|
"/v1/chat/completions",
|
||||||
|
strings.NewReader(`{"metadata":{"work_id":"42","scene":"suggestion_generation","muse_user_id":"u-1"}}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set(museHeaderUserID, "u-1")
|
||||||
|
req.Header.Set(museHeaderWorkID, "42")
|
||||||
|
req.Header.Set(museHeaderRequest, "req-1")
|
||||||
|
req.Header.Set(museHeaderTraceID, "trace-1")
|
||||||
|
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(MuseRequestContext())
|
||||||
|
router.POST("/v1/chat/completions", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, museRequestContextResponse{
|
||||||
|
MuseUserID: common.GetContextKeyString(c, constant.ContextKeyMuseUserID),
|
||||||
|
MuseWorkID: common.GetContextKeyString(c, constant.ContextKeyMuseWorkID),
|
||||||
|
MuseScene: common.GetContextKeyString(c, constant.ContextKeyMuseScene),
|
||||||
|
MuseRequestID: common.GetContextKeyString(c, constant.ContextKeyMuseRequestID),
|
||||||
|
TraceID: common.GetContextKeyString(c, constant.ContextKeyMuseTraceID),
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
router.ServeHTTP(recorder, req)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, recorder.Code)
|
||||||
|
|
||||||
|
var response museRequestContextResponse
|
||||||
|
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
|
||||||
|
require.Equal(t, "u-1", response.MuseUserID)
|
||||||
|
require.Equal(t, "42", response.MuseWorkID)
|
||||||
|
require.Equal(t, "suggestion_generation", response.MuseScene)
|
||||||
|
require.Equal(t, "trace-1", response.TraceID)
|
||||||
|
require.Equal(t, "req-1", response.MuseRequestID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMuseRequestContextFallsBackToMetadata(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPost,
|
||||||
|
"/v1/chat/completions",
|
||||||
|
strings.NewReader(`{"metadata":{"muse_user_id":"u-2","work_id":42,"scene":"suggestion_generation","request_id":"req-meta","trace_id":"trace-meta"}}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(MuseRequestContext())
|
||||||
|
router.POST("/v1/chat/completions", func(c *gin.Context) {
|
||||||
|
c.JSON(http.StatusOK, museRequestContextResponse{
|
||||||
|
MuseUserID: common.GetContextKeyString(c, constant.ContextKeyMuseUserID),
|
||||||
|
MuseWorkID: common.GetContextKeyString(c, constant.ContextKeyMuseWorkID),
|
||||||
|
MuseScene: common.GetContextKeyString(c, constant.ContextKeyMuseScene),
|
||||||
|
MuseRequestID: common.GetContextKeyString(c, constant.ContextKeyMuseRequestID),
|
||||||
|
TraceID: common.GetContextKeyString(c, constant.ContextKeyMuseTraceID),
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
router.ServeHTTP(recorder, req)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, recorder.Code)
|
||||||
|
|
||||||
|
var response museRequestContextResponse
|
||||||
|
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
|
||||||
|
require.Equal(t, "u-2", response.MuseUserID)
|
||||||
|
require.Equal(t, "42", response.MuseWorkID)
|
||||||
|
require.Equal(t, "suggestion_generation", response.MuseScene)
|
||||||
|
require.Equal(t, "trace-meta", response.TraceID)
|
||||||
|
require.Equal(t, "req-meta", response.MuseRequestID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMuseRequestContextRecordConsumeLogIncludesMuseFields(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
setupMuseRequestContextTestDB(t)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(
|
||||||
|
http.MethodPost,
|
||||||
|
"/v1/chat/completions",
|
||||||
|
strings.NewReader(`{"metadata":{"scene":"suggestion_generation","request_id":"req-log","trace_id":"trace-log"}}`),
|
||||||
|
)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set(museHeaderUserID, "u-log")
|
||||||
|
req.Header.Set(museHeaderWorkID, "77")
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(MuseRequestContext())
|
||||||
|
router.POST("/v1/chat/completions", func(c *gin.Context) {
|
||||||
|
model.RecordConsumeLog(c, 1, model.RecordConsumeLogParams{
|
||||||
|
ModelName: "gpt-4o",
|
||||||
|
Content: "ok",
|
||||||
|
Group: "default",
|
||||||
|
Other: map[string]interface{}{"existing": true},
|
||||||
|
})
|
||||||
|
var logEntry model.Log
|
||||||
|
require.NoError(t, model.LOG_DB.Last(&logEntry).Error)
|
||||||
|
c.JSON(http.StatusOK, gin.H{
|
||||||
|
"request_id": logEntry.RequestId,
|
||||||
|
"other": logEntry.Other,
|
||||||
|
})
|
||||||
|
})
|
||||||
|
|
||||||
|
router.ServeHTTP(recorder, req)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, recorder.Code)
|
||||||
|
var response struct {
|
||||||
|
RequestID string `json:"request_id"`
|
||||||
|
Other string `json:"other"`
|
||||||
|
}
|
||||||
|
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response))
|
||||||
|
require.Empty(t, response.RequestID)
|
||||||
|
require.Contains(t, response.Other, `"muse_user_id":"u-log"`)
|
||||||
|
require.Contains(t, response.Other, `"muse_work_id":"77"`)
|
||||||
|
require.Contains(t, response.Other, `"muse_request_id":"req-log"`)
|
||||||
|
require.Contains(t, response.Other, `"muse_scene":"suggestion_generation"`)
|
||||||
|
require.Contains(t, response.Other, `"muse_trace_id":"trace-log"`)
|
||||||
|
}
|
||||||
28
model/log.go
28
model/log.go
@ -7,6 +7,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/QuantumNous/new-api/common"
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/QuantumNous/new-api/constant"
|
||||||
"github.com/QuantumNous/new-api/logger"
|
"github.com/QuantumNous/new-api/logger"
|
||||||
"github.com/QuantumNous/new-api/types"
|
"github.com/QuantumNous/new-api/types"
|
||||||
|
|
||||||
@ -95,6 +96,7 @@ func RecordErrorLog(c *gin.Context, userId int, channelId int, modelName string,
|
|||||||
logger.LogInfo(c, fmt.Sprintf("record error log: userId=%d, channelId=%d, modelName=%s, tokenName=%s, content=%s", userId, channelId, modelName, tokenName, content))
|
logger.LogInfo(c, fmt.Sprintf("record error log: userId=%d, channelId=%d, modelName=%s, tokenName=%s, content=%s", userId, channelId, modelName, tokenName, content))
|
||||||
username := c.GetString("username")
|
username := c.GetString("username")
|
||||||
requestId := c.GetString(common.RequestIdKey)
|
requestId := c.GetString(common.RequestIdKey)
|
||||||
|
other = appendMuseContextToOther(c, other)
|
||||||
otherStr := common.MapToJsonStr(other)
|
otherStr := common.MapToJsonStr(other)
|
||||||
// 判断是否需要记录 IP
|
// 判断是否需要记录 IP
|
||||||
needRecordIp := false
|
needRecordIp := false
|
||||||
@ -156,6 +158,7 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams)
|
|||||||
logger.LogInfo(c, fmt.Sprintf("record consume log: userId=%d, params=%s", userId, common.GetJsonString(params)))
|
logger.LogInfo(c, fmt.Sprintf("record consume log: userId=%d, params=%s", userId, common.GetJsonString(params)))
|
||||||
username := c.GetString("username")
|
username := c.GetString("username")
|
||||||
requestId := c.GetString(common.RequestIdKey)
|
requestId := c.GetString(common.RequestIdKey)
|
||||||
|
params.Other = appendMuseContextToOther(c, params.Other)
|
||||||
otherStr := common.MapToJsonStr(params.Other)
|
otherStr := common.MapToJsonStr(params.Other)
|
||||||
// 判断是否需要记录 IP
|
// 判断是否需要记录 IP
|
||||||
needRecordIp := false
|
needRecordIp := false
|
||||||
@ -200,6 +203,31 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func appendMuseContextToOther(c *gin.Context, other map[string]interface{}) map[string]interface{} {
|
||||||
|
if c == nil {
|
||||||
|
return other
|
||||||
|
}
|
||||||
|
if other == nil {
|
||||||
|
other = make(map[string]interface{})
|
||||||
|
}
|
||||||
|
if museUserID := common.GetContextKeyString(c, constant.ContextKeyMuseUserID); museUserID != "" {
|
||||||
|
other["muse_user_id"] = museUserID
|
||||||
|
}
|
||||||
|
if museWorkID := common.GetContextKeyString(c, constant.ContextKeyMuseWorkID); museWorkID != "" {
|
||||||
|
other["muse_work_id"] = museWorkID
|
||||||
|
}
|
||||||
|
if museRequestID := common.GetContextKeyString(c, constant.ContextKeyMuseRequestID); museRequestID != "" {
|
||||||
|
other["muse_request_id"] = museRequestID
|
||||||
|
}
|
||||||
|
if museScene := common.GetContextKeyString(c, constant.ContextKeyMuseScene); museScene != "" {
|
||||||
|
other["muse_scene"] = museScene
|
||||||
|
}
|
||||||
|
if museTraceID := common.GetContextKeyString(c, constant.ContextKeyMuseTraceID); museTraceID != "" {
|
||||||
|
other["muse_trace_id"] = museTraceID
|
||||||
|
}
|
||||||
|
return other
|
||||||
|
}
|
||||||
|
|
||||||
type RecordTaskBillingLogParams struct {
|
type RecordTaskBillingLogParams struct {
|
||||||
UserId int
|
UserId int
|
||||||
LogType int
|
LogType int
|
||||||
|
|||||||
@ -139,7 +139,12 @@ type RelayInfo struct {
|
|||||||
SubscriptionPlanId int
|
SubscriptionPlanId int
|
||||||
SubscriptionPlanTitle string
|
SubscriptionPlanTitle string
|
||||||
// RequestId is used for idempotent pre-consume/refund
|
// RequestId is used for idempotent pre-consume/refund
|
||||||
RequestId string
|
RequestId string
|
||||||
|
MuseUserID string
|
||||||
|
MuseWorkID string
|
||||||
|
MuseRequestID string
|
||||||
|
MuseScene string
|
||||||
|
MuseTraceID string
|
||||||
// SubscriptionAmountTotal / SubscriptionAmountUsedAfterPreConsume are used to compute remaining in logs.
|
// SubscriptionAmountTotal / SubscriptionAmountUsedAfterPreConsume are used to compute remaining in logs.
|
||||||
SubscriptionAmountTotal int64
|
SubscriptionAmountTotal int64
|
||||||
SubscriptionAmountUsedAfterPreConsume int64
|
SubscriptionAmountUsedAfterPreConsume int64
|
||||||
@ -255,7 +260,6 @@ func (info *RelayInfo) ToString() string {
|
|||||||
latencyMs := info.FirstResponseTime.Sub(info.StartTime).Milliseconds()
|
latencyMs := info.FirstResponseTime.Sub(info.StartTime).Milliseconds()
|
||||||
fmt.Fprintf(b, "Timing{ Start: %s, FirstResponse: %s, LatencyMs: %d }, ",
|
fmt.Fprintf(b, "Timing{ Start: %s, FirstResponse: %s, LatencyMs: %d }, ",
|
||||||
info.StartTime.Format(time.RFC3339Nano), info.FirstResponseTime.Format(time.RFC3339Nano), latencyMs)
|
info.StartTime.Format(time.RFC3339Nano), info.FirstResponseTime.Format(time.RFC3339Nano), latencyMs)
|
||||||
|
|
||||||
// Audio / realtime
|
// Audio / realtime
|
||||||
if info.InputAudioFormat != "" || info.OutputAudioFormat != "" || len(info.RealtimeTools) > 0 || info.AudioUsage {
|
if info.InputAudioFormat != "" || info.OutputAudioFormat != "" || len(info.RealtimeTools) > 0 || info.AudioUsage {
|
||||||
fmt.Fprintf(b, "Realtime{ AudioUsage: %t, InFmt: %q, OutFmt: %q, Tools: %d }, ",
|
fmt.Fprintf(b, "Realtime{ AudioUsage: %t, InFmt: %q, OutFmt: %q, Tools: %d }, ",
|
||||||
@ -447,12 +451,17 @@ func genBaseRelayInfo(c *gin.Context, request dto.Request) *RelayInfo {
|
|||||||
info := &RelayInfo{
|
info := &RelayInfo{
|
||||||
Request: request,
|
Request: request,
|
||||||
|
|
||||||
RequestId: reqId,
|
RequestId: reqId,
|
||||||
UserId: common.GetContextKeyInt(c, constant.ContextKeyUserId),
|
MuseUserID: common.GetContextKeyString(c, constant.ContextKeyMuseUserID),
|
||||||
UsingGroup: common.GetContextKeyString(c, constant.ContextKeyUsingGroup),
|
MuseWorkID: common.GetContextKeyString(c, constant.ContextKeyMuseWorkID),
|
||||||
UserGroup: common.GetContextKeyString(c, constant.ContextKeyUserGroup),
|
MuseRequestID: common.GetContextKeyString(c, constant.ContextKeyMuseRequestID),
|
||||||
UserQuota: common.GetContextKeyInt(c, constant.ContextKeyUserQuota),
|
MuseScene: common.GetContextKeyString(c, constant.ContextKeyMuseScene),
|
||||||
UserEmail: common.GetContextKeyString(c, constant.ContextKeyUserEmail),
|
MuseTraceID: common.GetContextKeyString(c, constant.ContextKeyMuseTraceID),
|
||||||
|
UserId: common.GetContextKeyInt(c, constant.ContextKeyUserId),
|
||||||
|
UsingGroup: common.GetContextKeyString(c, constant.ContextKeyUsingGroup),
|
||||||
|
UserGroup: common.GetContextKeyString(c, constant.ContextKeyUserGroup),
|
||||||
|
UserQuota: common.GetContextKeyInt(c, constant.ContextKeyUserQuota),
|
||||||
|
UserEmail: common.GetContextKeyString(c, constant.ContextKeyUserEmail),
|
||||||
|
|
||||||
OriginModelName: common.GetContextKeyString(c, constant.ContextKeyOriginalModel),
|
OriginModelName: common.GetContextKeyString(c, constant.ContextKeyOriginalModel),
|
||||||
|
|
||||||
|
|||||||
@ -1,9 +1,15 @@
|
|||||||
package common
|
package common
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/QuantumNous/new-api/constant"
|
||||||
|
"github.com/QuantumNous/new-api/dto"
|
||||||
"github.com/QuantumNous/new-api/types"
|
"github.com/QuantumNous/new-api/types"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -38,3 +44,30 @@ func TestRelayInfoGetFinalRequestRelayFormatNilReceiver(t *testing.T) {
|
|||||||
var info *RelayInfo
|
var info *RelayInfo
|
||||||
require.Equal(t, types.RelayFormat(""), info.GetFinalRequestRelayFormat())
|
require.Equal(t, types.RelayFormat(""), info.GetFinalRequestRelayFormat())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMuseRequestContextGenRelayInfoCopiesMuseFields(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
ctx, _ := gin.CreateTestContext(recorder)
|
||||||
|
ctx.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
|
||||||
|
|
||||||
|
ctx.Set(common.RequestIdKey, "req-1")
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyUserId, 11)
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyUserGroup, "default")
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyUsingGroup, "default")
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyMuseUserID, "muse-user-1")
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyMuseWorkID, "42")
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyMuseRequestID, "muse-req-1")
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyMuseScene, "suggestion_generation")
|
||||||
|
common.SetContextKey(ctx, constant.ContextKeyMuseTraceID, "trace-1")
|
||||||
|
|
||||||
|
info := GenRelayInfoOpenAI(ctx, &dto.GeneralOpenAIRequest{})
|
||||||
|
|
||||||
|
require.NotNil(t, info)
|
||||||
|
require.Equal(t, "req-1", info.RequestId)
|
||||||
|
require.Equal(t, "muse-user-1", info.MuseUserID)
|
||||||
|
require.Equal(t, "42", info.MuseWorkID)
|
||||||
|
require.Equal(t, "muse-req-1", info.MuseRequestID)
|
||||||
|
require.Equal(t, "suggestion_generation", info.MuseScene)
|
||||||
|
require.Equal(t, "trace-1", info.MuseTraceID)
|
||||||
|
}
|
||||||
|
|||||||
@ -53,6 +53,15 @@ func SetApiRouter(router *gin.Engine) {
|
|||||||
// Universal secure verification routes
|
// Universal secure verification routes
|
||||||
apiRouter.POST("/verify", middleware.UserAuth(), middleware.CriticalRateLimit(), controller.UniversalVerify)
|
apiRouter.POST("/verify", middleware.UserAuth(), middleware.CriticalRateLimit(), controller.UniversalVerify)
|
||||||
|
|
||||||
|
museInternalRoute := apiRouter.Group("/internal/muse")
|
||||||
|
museInternalRoute.Use(middleware.MuseInternalAuth())
|
||||||
|
{
|
||||||
|
museInternalRoute.POST("/users/sync", controller.SyncMuseUser)
|
||||||
|
museInternalRoute.POST("/users/:id/disable", controller.DisableMuseUser)
|
||||||
|
museInternalRoute.POST("/tokens/:id/revoke", controller.RevokeMuseToken)
|
||||||
|
museInternalRoute.GET("/users/:id/status", controller.GetMuseUserStatus)
|
||||||
|
}
|
||||||
|
|
||||||
userRoute := apiRouter.Group("/user")
|
userRoute := apiRouter.Group("/user")
|
||||||
{
|
{
|
||||||
userRoute.POST("/register", middleware.CriticalRateLimit(), middleware.TurnstileCheck(), controller.Register)
|
userRoute.POST("/register", middleware.CriticalRateLimit(), middleware.TurnstileCheck(), controller.Register)
|
||||||
|
|||||||
@ -70,6 +70,7 @@ func SetRelayRouter(router *gin.Engine) {
|
|||||||
relayV1Router.Use(middleware.RouteTag("relay"))
|
relayV1Router.Use(middleware.RouteTag("relay"))
|
||||||
relayV1Router.Use(middleware.SystemPerformanceCheck())
|
relayV1Router.Use(middleware.SystemPerformanceCheck())
|
||||||
relayV1Router.Use(middleware.TokenAuth())
|
relayV1Router.Use(middleware.TokenAuth())
|
||||||
|
relayV1Router.Use(middleware.MuseRequestContext())
|
||||||
relayV1Router.Use(middleware.ModelRequestRateLimit())
|
relayV1Router.Use(middleware.ModelRequestRateLimit())
|
||||||
{
|
{
|
||||||
// WebSocket 路由(统一到 Relay)
|
// WebSocket 路由(统一到 Relay)
|
||||||
@ -190,6 +191,7 @@ func SetRelayRouter(router *gin.Engine) {
|
|||||||
relayGeminiRouter.Use(middleware.RouteTag("relay"))
|
relayGeminiRouter.Use(middleware.RouteTag("relay"))
|
||||||
relayGeminiRouter.Use(middleware.SystemPerformanceCheck())
|
relayGeminiRouter.Use(middleware.SystemPerformanceCheck())
|
||||||
relayGeminiRouter.Use(middleware.TokenAuth())
|
relayGeminiRouter.Use(middleware.TokenAuth())
|
||||||
|
relayGeminiRouter.Use(middleware.MuseRequestContext())
|
||||||
relayGeminiRouter.Use(middleware.ModelRequestRateLimit())
|
relayGeminiRouter.Use(middleware.ModelRequestRateLimit())
|
||||||
relayGeminiRouter.Use(middleware.Distribute())
|
relayGeminiRouter.Use(middleware.Distribute())
|
||||||
{
|
{
|
||||||
|
|||||||
@ -90,7 +90,7 @@ func CacheGetRandomSatisfiedChannel(param *RetryParam) (*model.Channel, string,
|
|||||||
if len(setting.GetAutoGroups()) == 0 {
|
if len(setting.GetAutoGroups()) == 0 {
|
||||||
return nil, selectGroup, errors.New("auto groups is not enabled")
|
return nil, selectGroup, errors.New("auto groups is not enabled")
|
||||||
}
|
}
|
||||||
autoGroups := GetUserAutoGroup(userGroup)
|
autoGroups := ResolveMusePreferredGroups(param.Ctx, userGroup)
|
||||||
|
|
||||||
// startGroupIndex: the group index to start searching from
|
// startGroupIndex: the group index to start searching from
|
||||||
// startGroupIndex: 开始搜索的分组索引
|
// startGroupIndex: 开始搜索的分组索引
|
||||||
|
|||||||
305
service/muse_internal_service.go
Normal file
305
service/muse_internal_service.go
Normal file
@ -0,0 +1,305 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
museDTO "github.com/QuantumNous/new-api/dto"
|
||||||
|
"github.com/QuantumNous/new-api/model"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
museUserStatusActive = "active"
|
||||||
|
museUserStatusDisabled = "disabled"
|
||||||
|
museMirrorTokenName = "muse-long-lived"
|
||||||
|
museRemarkPrefix = "muse_user_id:"
|
||||||
|
)
|
||||||
|
|
||||||
|
func SyncMuseUser(req museDTO.SyncMuseUserRequest) (*museDTO.SyncMuseUserResponse, error) {
|
||||||
|
req.MuseUserID = strings.TrimSpace(req.MuseUserID)
|
||||||
|
req.Username = strings.TrimSpace(req.Username)
|
||||||
|
req.DisplayName = strings.TrimSpace(req.DisplayName)
|
||||||
|
if req.MuseUserID == "" {
|
||||||
|
return nil, errors.New("muse_user_id is required")
|
||||||
|
}
|
||||||
|
|
||||||
|
if normalizeMuseStatus(req.Status) == museUserStatusDisabled {
|
||||||
|
return disableMuseUserByMuseUserID(req.MuseUserID)
|
||||||
|
}
|
||||||
|
|
||||||
|
user, err := getMirrorUserByMuseUserID(req.MuseUserID)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
user, err = createMirrorUser(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
} else if err = updateMirrorUser(user, req); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
token, err := ensureMirrorToken(user.Id)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return buildMirrorResponse(user, token, museUserStatusActive, true), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func DisableMuseUserByID(newAPIUserID int) (*museDTO.SyncMuseUserResponse, error) {
|
||||||
|
if newAPIUserID <= 0 {
|
||||||
|
return nil, errors.New("newapi user id is required")
|
||||||
|
}
|
||||||
|
user, err := model.GetUserById(newAPIUserID, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err = disableUserAndTokens(user); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
token, _ := getMirrorToken(user.Id)
|
||||||
|
return buildMirrorResponse(user, token, museUserStatusDisabled, false), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func RevokeMuseToken(tokenID int) (*museDTO.SyncMuseUserResponse, error) {
|
||||||
|
if tokenID <= 0 {
|
||||||
|
return nil, errors.New("token id is required")
|
||||||
|
}
|
||||||
|
token, err := model.GetTokenById(tokenID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
token.Status = common.TokenStatusDisabled
|
||||||
|
if err = token.Update(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
user, err := model.GetUserById(token.UserId, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return buildMirrorResponse(user, token, museUserStatusDisabled, false), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetMuseUserStatus(newAPIUserID int) (*museDTO.SyncMuseUserResponse, error) {
|
||||||
|
if newAPIUserID <= 0 {
|
||||||
|
return nil, errors.New("newapi user id is required")
|
||||||
|
}
|
||||||
|
user, err := model.GetUserById(newAPIUserID, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
token, _ := getMirrorToken(user.Id)
|
||||||
|
return buildMirrorResponse(user, token, userStatusToMuseStatus(user.Status), false), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func disableMuseUserByMuseUserID(museUserID string) (*museDTO.SyncMuseUserResponse, error) {
|
||||||
|
user, err := getMirrorUserByMuseUserID(museUserID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err = disableUserAndTokens(user); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
token, _ := getMirrorToken(user.Id)
|
||||||
|
return buildMirrorResponse(user, token, museUserStatusDisabled, false), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func getMirrorUserByMuseUserID(museUserID string) (*model.User, error) {
|
||||||
|
user := &model.User{}
|
||||||
|
err := model.DB.Where("remark = ?", museRemarkPrefix+museUserID).First(user).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func createMirrorUser(req museDTO.SyncMuseUserRequest) (*model.User, error) {
|
||||||
|
hashedPassword, err := common.Password2Hash(common.GetRandomString(20))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
user := &model.User{
|
||||||
|
Username: chooseMirrorUsername(req.Username, req.MuseUserID, 0, ""),
|
||||||
|
Password: hashedPassword,
|
||||||
|
DisplayName: chooseDisplayName(req.DisplayName, req.Username),
|
||||||
|
Role: common.RoleCommonUser,
|
||||||
|
Status: common.UserStatusEnabled,
|
||||||
|
Group: "default",
|
||||||
|
Remark: museRemarkPrefix + req.MuseUserID,
|
||||||
|
}
|
||||||
|
if user.DisplayName == "" {
|
||||||
|
user.DisplayName = user.Username
|
||||||
|
}
|
||||||
|
if err = model.DB.Create(user).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func updateMirrorUser(user *model.User, req museDTO.SyncMuseUserRequest) error {
|
||||||
|
username := chooseMirrorUsername(req.Username, req.MuseUserID, user.Id, user.Username)
|
||||||
|
displayName := chooseDisplayName(req.DisplayName, req.Username)
|
||||||
|
if displayName == "" {
|
||||||
|
displayName = user.DisplayName
|
||||||
|
}
|
||||||
|
if displayName == "" {
|
||||||
|
displayName = username
|
||||||
|
}
|
||||||
|
|
||||||
|
user.Username = username
|
||||||
|
user.DisplayName = displayName
|
||||||
|
user.Status = common.UserStatusEnabled
|
||||||
|
user.Role = common.RoleCommonUser
|
||||||
|
user.Remark = museRemarkPrefix + req.MuseUserID
|
||||||
|
return user.Update(false)
|
||||||
|
}
|
||||||
|
|
||||||
|
func ensureMirrorToken(userID int) (*model.Token, error) {
|
||||||
|
token, err := getMirrorToken(userID)
|
||||||
|
if err != nil {
|
||||||
|
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
key, keyErr := common.GenerateKey()
|
||||||
|
if keyErr != nil {
|
||||||
|
return nil, keyErr
|
||||||
|
}
|
||||||
|
token = &model.Token{
|
||||||
|
UserId: userID,
|
||||||
|
Name: museMirrorTokenName,
|
||||||
|
Key: key,
|
||||||
|
Status: common.TokenStatusEnabled,
|
||||||
|
CreatedTime: common.GetTimestamp(),
|
||||||
|
AccessedTime: common.GetTimestamp(),
|
||||||
|
ExpiredTime: -1,
|
||||||
|
RemainQuota: 0,
|
||||||
|
UnlimitedQuota: true,
|
||||||
|
}
|
||||||
|
if err = token.Insert(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return token, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
token.Status = common.TokenStatusEnabled
|
||||||
|
token.ExpiredTime = -1
|
||||||
|
token.UnlimitedQuota = true
|
||||||
|
token.RemainQuota = 0
|
||||||
|
token.AccessedTime = common.GetTimestamp()
|
||||||
|
if err = token.Update(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return token, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func getMirrorToken(userID int) (*model.Token, error) {
|
||||||
|
token := &model.Token{}
|
||||||
|
err := model.DB.Where("user_id = ? AND name = ?", userID, museMirrorTokenName).Order("id desc").First(token).Error
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return token, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func disableUserAndTokens(user *model.User) error {
|
||||||
|
user.Status = common.UserStatusDisabled
|
||||||
|
if err := user.Update(false); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
tokens := make([]model.Token, 0)
|
||||||
|
if err := model.DB.Where("user_id = ?", user.Id).Find(&tokens).Error; err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for i := range tokens {
|
||||||
|
tokens[i].Status = common.TokenStatusDisabled
|
||||||
|
if err := tokens[i].Update(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildMirrorResponse(user *model.User, token *model.Token, status string, includeTokenKey bool) *museDTO.SyncMuseUserResponse {
|
||||||
|
response := &museDTO.SyncMuseUserResponse{Status: status}
|
||||||
|
if user != nil {
|
||||||
|
response.NewAPIUserID = user.Id
|
||||||
|
}
|
||||||
|
if token != nil {
|
||||||
|
response.TokenID = token.Id
|
||||||
|
if includeTokenKey {
|
||||||
|
response.TokenKey = token.GetFullKey()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return response
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMuseStatus(status string) string {
|
||||||
|
if strings.EqualFold(strings.TrimSpace(status), museUserStatusDisabled) {
|
||||||
|
return museUserStatusDisabled
|
||||||
|
}
|
||||||
|
return museUserStatusActive
|
||||||
|
}
|
||||||
|
|
||||||
|
func chooseDisplayName(displayName string, username string) string {
|
||||||
|
if strings.TrimSpace(displayName) != "" {
|
||||||
|
return trimToLength(displayName, model.UserNameMaxLength)
|
||||||
|
}
|
||||||
|
return trimToLength(username, model.UserNameMaxLength)
|
||||||
|
}
|
||||||
|
|
||||||
|
func chooseMirrorUsername(desired string, museUserID string, currentUserID int, currentUsername string) string {
|
||||||
|
desired = trimToLength(desired, model.UserNameMaxLength)
|
||||||
|
if desired != "" && usernameAvailable(desired, currentUserID) {
|
||||||
|
return desired
|
||||||
|
}
|
||||||
|
currentUsername = trimToLength(currentUsername, model.UserNameMaxLength)
|
||||||
|
if currentUsername != "" && usernameAvailable(currentUsername, currentUserID) {
|
||||||
|
return currentUsername
|
||||||
|
}
|
||||||
|
fallback := trimToLength(fmt.Sprintf("muse_%s", common.GenerateHMAC(museUserID)[:15]), model.UserNameMaxLength)
|
||||||
|
if usernameAvailable(fallback, currentUserID) {
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
candidate := trimToLength("m"+common.GetRandomString(model.UserNameMaxLength-1), model.UserNameMaxLength)
|
||||||
|
if usernameAvailable(candidate, currentUserID) {
|
||||||
|
return candidate
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
|
||||||
|
func usernameAvailable(username string, currentUserID int) bool {
|
||||||
|
if username == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
query := model.DB.Model(&model.User{}).Where("username = ?", username)
|
||||||
|
if currentUserID > 0 {
|
||||||
|
query = query.Where("id <> ?", currentUserID)
|
||||||
|
}
|
||||||
|
var count int64
|
||||||
|
if err := query.Count(&count).Error; err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return count == 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func trimToLength(value string, limit int) string {
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
if len(value) <= limit {
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
return value[:limit]
|
||||||
|
}
|
||||||
|
|
||||||
|
func userStatusToMuseStatus(status int) string {
|
||||||
|
if status == common.UserStatusDisabled {
|
||||||
|
return museUserStatusDisabled
|
||||||
|
}
|
||||||
|
return museUserStatusActive
|
||||||
|
}
|
||||||
146
service/muse_routing_policy.go
Normal file
146
service/muse_routing_policy.go
Normal file
@ -0,0 +1,146 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/QuantumNous/new-api/common"
|
||||||
|
"github.com/QuantumNous/new-api/constant"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
const ginKeyMuseRoutingPolicy = "muse_routing_policy"
|
||||||
|
|
||||||
|
type MuseRoutingPolicy struct {
|
||||||
|
Scene string
|
||||||
|
AllowCrossGroupRetry bool
|
||||||
|
AllowModelDegrade bool
|
||||||
|
PreferredGroups []string
|
||||||
|
}
|
||||||
|
|
||||||
|
func ResolveMuseRoutingPolicy(scene string, routingPolicy string, degradePolicy string) MuseRoutingPolicy {
|
||||||
|
scene = normalizeMusePolicyValue(scene)
|
||||||
|
routingPolicy = normalizeMusePolicyValue(routingPolicy)
|
||||||
|
degradePolicy = normalizeMusePolicyValue(degradePolicy)
|
||||||
|
|
||||||
|
policy := MuseRoutingPolicy{
|
||||||
|
Scene: scene,
|
||||||
|
PreferredGroups: defaultMusePreferredGroups(scene, routingPolicy),
|
||||||
|
}
|
||||||
|
|
||||||
|
switch routingPolicy {
|
||||||
|
case "quality_first", "balanced":
|
||||||
|
policy.AllowCrossGroupRetry = true
|
||||||
|
}
|
||||||
|
|
||||||
|
if degradePolicy == "allow_lower_tier" {
|
||||||
|
policy.AllowModelDegrade = true
|
||||||
|
}
|
||||||
|
|
||||||
|
return policy
|
||||||
|
}
|
||||||
|
|
||||||
|
func SetMuseRoutingPolicy(c *gin.Context, policy MuseRoutingPolicy) {
|
||||||
|
if c == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if policy.Scene == "" && !policy.AllowCrossGroupRetry && !policy.AllowModelDegrade && len(policy.PreferredGroups) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Set(ginKeyMuseRoutingPolicy, policy)
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetMuseRoutingPolicy(c *gin.Context) (MuseRoutingPolicy, bool) {
|
||||||
|
if c == nil {
|
||||||
|
return MuseRoutingPolicy{}, false
|
||||||
|
}
|
||||||
|
value, ok := c.Get(ginKeyMuseRoutingPolicy)
|
||||||
|
if !ok {
|
||||||
|
return MuseRoutingPolicy{}, false
|
||||||
|
}
|
||||||
|
policy, ok := value.(MuseRoutingPolicy)
|
||||||
|
if !ok {
|
||||||
|
return MuseRoutingPolicy{}, false
|
||||||
|
}
|
||||||
|
return policy, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func ApplyMuseRoutingPolicy(c *gin.Context, usingGroup string) string {
|
||||||
|
policy, ok := GetMuseRoutingPolicy(c)
|
||||||
|
if !ok {
|
||||||
|
return usingGroup
|
||||||
|
}
|
||||||
|
if policy.AllowCrossGroupRetry {
|
||||||
|
common.SetContextKey(c, constant.ContextKeyTokenCrossGroupRetry, true)
|
||||||
|
if usingGroup != "auto" {
|
||||||
|
usingGroup = "auto"
|
||||||
|
common.SetContextKey(c, constant.ContextKeyUsingGroup, usingGroup)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return usingGroup
|
||||||
|
}
|
||||||
|
|
||||||
|
func ResolveMusePreferredGroups(c *gin.Context, userGroup string) []string {
|
||||||
|
policy, ok := GetMuseRoutingPolicy(c)
|
||||||
|
if !ok || len(policy.PreferredGroups) == 0 {
|
||||||
|
return GetUserAutoGroup(userGroup)
|
||||||
|
}
|
||||||
|
usableGroups := GetUserUsableGroups(userGroup)
|
||||||
|
filtered := make([]string, 0, len(policy.PreferredGroups))
|
||||||
|
for _, group := range policy.PreferredGroups {
|
||||||
|
if _, allowed := usableGroups[group]; allowed {
|
||||||
|
filtered = append(filtered, group)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(filtered) == 0 {
|
||||||
|
return GetUserAutoGroup(userGroup)
|
||||||
|
}
|
||||||
|
return filtered
|
||||||
|
}
|
||||||
|
|
||||||
|
func MuseFallbackModelCandidates(modelName string, policy MuseRoutingPolicy) []string {
|
||||||
|
candidates := []string{modelName}
|
||||||
|
if !policy.AllowModelDegrade {
|
||||||
|
return candidates
|
||||||
|
}
|
||||||
|
|
||||||
|
appendUnique := func(candidate string) {
|
||||||
|
candidate = strings.TrimSpace(candidate)
|
||||||
|
if candidate == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !slices.Contains(candidates, candidate) {
|
||||||
|
candidates = append(candidates, candidate)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switch {
|
||||||
|
case strings.HasPrefix(modelName, "gpt-4o") && modelName != "gpt-4o-mini":
|
||||||
|
appendUnique("gpt-4o-mini")
|
||||||
|
case strings.Contains(modelName, "gemini") && strings.Contains(modelName, "pro"):
|
||||||
|
appendUnique(strings.Replace(modelName, "pro", "flash", 1))
|
||||||
|
case strings.Contains(modelName, "claude") && strings.Contains(modelName, "opus"):
|
||||||
|
appendUnique(strings.Replace(modelName, "opus", "sonnet", 1))
|
||||||
|
case strings.Contains(modelName, "claude") && strings.Contains(modelName, "sonnet"):
|
||||||
|
appendUnique(strings.Replace(modelName, "sonnet", "haiku", 1))
|
||||||
|
}
|
||||||
|
|
||||||
|
return candidates
|
||||||
|
}
|
||||||
|
|
||||||
|
func defaultMusePreferredGroups(scene string, routingPolicy string) []string {
|
||||||
|
switch routingPolicy {
|
||||||
|
case "quality_first":
|
||||||
|
return []string{"svip", "vip", "default"}
|
||||||
|
case "balanced":
|
||||||
|
return []string{"vip", "default"}
|
||||||
|
}
|
||||||
|
if scene == "suggestion_generation" {
|
||||||
|
return []string{"default"}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeMusePolicyValue(value string) string {
|
||||||
|
return strings.TrimSpace(strings.ToLower(value))
|
||||||
|
}
|
||||||
43
service/muse_routing_policy_test.go
Normal file
43
service/muse_routing_policy_test.go
Normal file
@ -0,0 +1,43 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestResolveMuseRoutingPolicy_QualityFirstAllowsFallback(t *testing.T) {
|
||||||
|
policy := ResolveMuseRoutingPolicy("suggestion_generation", "quality_first", "allow_lower_tier")
|
||||||
|
|
||||||
|
require.Equal(t, "suggestion_generation", policy.Scene)
|
||||||
|
require.True(t, policy.AllowCrossGroupRetry)
|
||||||
|
require.True(t, policy.AllowModelDegrade)
|
||||||
|
require.Equal(t, []string{"svip", "vip", "default"}, policy.PreferredGroups)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveMuseRoutingPolicy_DefaultsWithoutPolicies(t *testing.T) {
|
||||||
|
policy := ResolveMuseRoutingPolicy("suggestion_generation", "", "")
|
||||||
|
|
||||||
|
require.False(t, policy.AllowCrossGroupRetry)
|
||||||
|
require.False(t, policy.AllowModelDegrade)
|
||||||
|
require.Equal(t, []string{"default"}, policy.PreferredGroups)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMuseFallbackModelCandidates_AllowLowerTierAddsFallback(t *testing.T) {
|
||||||
|
policy := ResolveMuseRoutingPolicy("suggestion_generation", "quality_first", "allow_lower_tier")
|
||||||
|
|
||||||
|
candidates := MuseFallbackModelCandidates("gpt-4o", policy)
|
||||||
|
|
||||||
|
require.Equal(t, []string{"gpt-4o", "gpt-4o-mini"}, candidates)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveMusePreferredGroups_UsesSceneDrivenDefaultsFromContext(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
ctx, _ := gin.CreateTestContext(nil)
|
||||||
|
SetMuseRoutingPolicy(ctx, ResolveMuseRoutingPolicy("suggestion_generation", "", ""))
|
||||||
|
|
||||||
|
groups := ResolveMusePreferredGroups(ctx, "default")
|
||||||
|
|
||||||
|
require.Equal(t, []string{"default"}, groups)
|
||||||
|
}
|
||||||
Loading…
x
Reference in New Issue
Block a user