mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-14 08:13:37 +00:00
feat(task): replace built-in task adaptors with a sandboxed JS plugin system (#7076)
This commit is contained in:
+212
-35
@@ -6,7 +6,6 @@ import (
|
||||
"io"
|
||||
"net/http"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -14,7 +13,9 @@ import (
|
||||
"github.com/QuantumNous/new-api/constant"
|
||||
taskdto "github.com/QuantumNous/new-api/dto"
|
||||
"github.com/QuantumNous/new-api/i18n"
|
||||
"github.com/QuantumNous/new-api/logger"
|
||||
"github.com/QuantumNous/new-api/model"
|
||||
"github.com/QuantumNous/new-api/pkg/jsplugin"
|
||||
relayconstant "github.com/QuantumNous/new-api/relay/constant"
|
||||
"github.com/QuantumNous/new-api/relaykit/dto"
|
||||
"github.com/QuantumNous/new-api/relaykit/types"
|
||||
@@ -33,25 +34,46 @@ type ModelRequest struct {
|
||||
func Distribute() func(c *gin.Context) {
|
||||
return func(c *gin.Context) {
|
||||
var channel *model.Channel
|
||||
channelId, ok := common.GetContextKey(c, constant.ContextKeyTokenSpecificChannelId)
|
||||
constraints := service.GetChannelConstraints(c)
|
||||
constraints.AddFilter(taskdto.ChannelFilter{
|
||||
Kind: taskdto.FilterRequestPath,
|
||||
RequestPath: c.Request.URL.Path,
|
||||
})
|
||||
service.AppendTaskPluginIdentityFilter(c, c.GetString("expected_task_plugin_key"))
|
||||
modelRequest, shouldSelectChannel, err := getModelRequest(c)
|
||||
if err != nil {
|
||||
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidRequest, map[string]any{"Error": err.Error()}))
|
||||
return
|
||||
}
|
||||
if ok {
|
||||
id, err := strconv.Atoi(channelId.(string))
|
||||
if err != nil {
|
||||
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidChannelId))
|
||||
return
|
||||
if pin, found, overridden := constraints.ResolvedPin(); found {
|
||||
for _, lost := range overridden {
|
||||
logger.LogWarn(c, fmt.Sprintf(
|
||||
"channel pin overridden: winning_source=%s winning_channel_id=%d overridden_source=%s overridden_channel_id=%d",
|
||||
pin.Source, pin.ChannelId, lost.Source, lost.ChannelId,
|
||||
))
|
||||
}
|
||||
channel, err = model.GetChannelById(id, true)
|
||||
channel, err = model.CacheGetChannel(pin.ChannelId)
|
||||
if err != nil {
|
||||
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidChannelId))
|
||||
if pin.Source == taskdto.PinSourceOriginTask {
|
||||
abortWithOpenAiMessage(c, http.StatusBadRequest, "origin_task_channel_disabled", types.ErrorCode("origin_task_channel_disabled"))
|
||||
} else {
|
||||
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorInvalidChannelId))
|
||||
}
|
||||
return
|
||||
}
|
||||
if channel.Status != common.ChannelStatusEnabled {
|
||||
abortWithOpenAiMessage(c, http.StatusForbidden, i18n.T(c, i18n.MsgDistributorChannelDisabled))
|
||||
if pin.Source == taskdto.PinSourceOriginTask {
|
||||
abortWithOpenAiMessage(c, http.StatusBadRequest, "origin_task_channel_disabled", types.ErrorCode("origin_task_channel_disabled"))
|
||||
} else {
|
||||
abortWithOpenAiMessage(c, http.StatusForbidden, i18n.T(c, i18n.MsgDistributorChannelDisabled))
|
||||
}
|
||||
return
|
||||
}
|
||||
if ok, kind := model.ChannelSatisfiesFilters(channel, modelRequest.Model, constraints.Filters); !ok {
|
||||
if kind == taskdto.FilterTaskPluginIdentity {
|
||||
logTaskPluginChannelDecision(c, channel, modelRequest.Model, "channel_rejected", "identity_mismatch")
|
||||
}
|
||||
abortWithOpenAiMessage(c, http.StatusBadRequest, i18n.T(c, i18n.MsgDistributorNoAvailableChannel, map[string]any{"Group": common.GetContextKeyString(c, constant.ContextKeyUsingGroup), "Model": modelRequest.Model}), types.ErrorCode(kind))
|
||||
return
|
||||
}
|
||||
} else {
|
||||
@@ -105,8 +127,11 @@ func Distribute() func(c *gin.Context) {
|
||||
if preferredChannelID, found := service.GetPreferredChannelByAffinity(c, modelRequest.Model, usingGroup); found {
|
||||
affinityUsable := false
|
||||
preferred, err := model.CacheGetChannel(preferredChannelID)
|
||||
if err == nil && preferred != nil && preferred.Status == common.ChannelStatusEnabled &&
|
||||
channelSupportsRequestPath(preferred, c.Request.URL.Path, modelRequest.Model) {
|
||||
affinitySatisfied := false
|
||||
if err == nil && preferred != nil && preferred.Status == common.ChannelStatusEnabled {
|
||||
affinitySatisfied, _ = model.ChannelSatisfiesFilters(preferred, modelRequest.Model, constraints.Filters)
|
||||
}
|
||||
if affinitySatisfied {
|
||||
if usingGroup == "auto" {
|
||||
userGroup := common.GetContextKeyString(c, constant.ContextKeyUserGroup)
|
||||
autoGroups := service.GetRequestAutoGroups(c, userGroup)
|
||||
@@ -161,6 +186,15 @@ func Distribute() func(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
}
|
||||
if channel != nil {
|
||||
if ok, kind := model.ChannelSatisfiesFilters(channel, modelRequest.Model, constraints.Filters); !ok {
|
||||
if kind == taskdto.FilterTaskPluginIdentity {
|
||||
logTaskPluginChannelDecision(c, channel, modelRequest.Model, "channel_rejected", "identity_mismatch")
|
||||
}
|
||||
abortWithOpenAiMessage(c, http.StatusServiceUnavailable, i18n.T(c, i18n.MsgDistributorNoAvailableChannel, map[string]any{"Group": common.GetContextKeyString(c, constant.ContextKeyUsingGroup), "Model": modelRequest.Model}), types.ErrorCodeModelNotFound)
|
||||
return
|
||||
}
|
||||
}
|
||||
common.SetContextKey(c, constant.ContextKeyRequestStartTime, time.Now())
|
||||
SetupContextForSelectedChannel(c, channel, modelRequest.Model)
|
||||
c.Next()
|
||||
@@ -170,18 +204,68 @@ func Distribute() func(c *gin.Context) {
|
||||
}
|
||||
}
|
||||
|
||||
// channelSupportsRequestPath reports whether a channel can serve the request path.
|
||||
// Only Advanced Custom (type 58) channels are path-checked; all other channel types
|
||||
// always pass. A type-58 channel is usable only when one of its routes matches.
|
||||
func channelSupportsRequestPath(channel *model.Channel, requestPath string, requestModel string) bool {
|
||||
func channelMatchesExpectedTaskPlugin(c *gin.Context, channel *model.Channel, expected string) bool {
|
||||
if channel == nil {
|
||||
return false
|
||||
}
|
||||
if channel.Type != constant.ChannelTypeAdvancedCustom {
|
||||
if c != nil {
|
||||
if _, matched := pinnedEndpointCandidateForChannel(c, channel, expected); matched {
|
||||
return true
|
||||
}
|
||||
}
|
||||
if channel.Type == constant.ChannelTypeTaskPlugin {
|
||||
return expected != "" && channel.GetSetting().TaskPluginKey == expected
|
||||
}
|
||||
if expected == "" {
|
||||
return true
|
||||
}
|
||||
config := channel.GetOtherSettings().AdvancedCustom
|
||||
return config != nil && config.SupportsPathForModel(requestPath, requestModel)
|
||||
|
||||
if c == nil {
|
||||
return false
|
||||
}
|
||||
value, exists := c.Get(jsplugin.ContextKeyPinnedPlugin)
|
||||
pinned, ok := value.(jsplugin.PinnedPlugin)
|
||||
if !exists || !ok || pinned.Generation == nil || pinned.Plugin == nil || pinned.Plugin.Meta.Key != expected {
|
||||
return false
|
||||
}
|
||||
plugin, ok := pinned.Generation.GetByChannelType(channel.Type)
|
||||
return ok && plugin == pinned.Plugin
|
||||
}
|
||||
|
||||
func pinnedEndpointCandidateForChannel(c *gin.Context, channel *model.Channel, expected string) (jsplugin.ProtocolBinding, bool) {
|
||||
if c == nil || channel == nil || expected == "" {
|
||||
return jsplugin.ProtocolBinding{}, false
|
||||
}
|
||||
value, exists := c.Get(jsplugin.ContextKeyPinnedEndpoint)
|
||||
pinned, ok := value.(jsplugin.PinnedEndpoint)
|
||||
if !exists || !ok || pinned.Generation == nil || pinned.Plugin == nil {
|
||||
return jsplugin.ProtocolBinding{}, false
|
||||
}
|
||||
candidates := pinned.Candidates
|
||||
if len(candidates) == 0 {
|
||||
candidates = []jsplugin.ProtocolBinding{{Plugin: pinned.Plugin, Protocol: pinned.Protocol, Operation: pinned.Operation, Model: pinned.Model}}
|
||||
}
|
||||
expectedOwned := false
|
||||
selected := jsplugin.ProtocolBinding{}
|
||||
for _, candidate := range candidates {
|
||||
if candidate.Plugin == nil {
|
||||
continue
|
||||
}
|
||||
if candidate.Plugin.Meta.Key == expected {
|
||||
expectedOwned = true
|
||||
}
|
||||
if channel.Type == constant.ChannelTypeTaskPlugin {
|
||||
if channel.GetSetting().TaskPluginKey == candidate.Plugin.Meta.Key {
|
||||
selected = candidate
|
||||
}
|
||||
continue
|
||||
}
|
||||
plugin, indexed := pinned.Generation.GetByChannelType(channel.Type)
|
||||
if indexed && plugin == candidate.Plugin {
|
||||
selected = candidate
|
||||
}
|
||||
}
|
||||
return selected, expectedOwned && selected.Plugin != nil
|
||||
}
|
||||
|
||||
// getModelFromRequest 从请求中读取模型信息
|
||||
@@ -190,6 +274,12 @@ func channelSupportsRequestPath(channel *model.Channel, requestPath string, requ
|
||||
// - application/x-www-form-urlencoded
|
||||
// - multipart/form-data
|
||||
func getModelFromRequest(c *gin.Context) (*ModelRequest, error) {
|
||||
if cached, exists := c.Get(contextKeyTaskPluginEndpointModel); exists {
|
||||
if modelRequest, ok := cached.(ModelRequest); ok {
|
||||
cachedRequest := modelRequest
|
||||
return &cachedRequest, nil
|
||||
}
|
||||
}
|
||||
if strings.HasPrefix(c.Request.Header.Get("Content-Type"), "application/json") {
|
||||
modelRequest, err := getModelFromJSONBody(c)
|
||||
if err != nil {
|
||||
@@ -218,6 +308,9 @@ func getModelFromJSONBody(c *gin.Context) (*ModelRequest, error) {
|
||||
if !gjson.ValidBytes(requestBody) {
|
||||
return nil, errors.New("invalid JSON request body")
|
||||
}
|
||||
if countTopLevelJSONKey(requestBody, "model") > 1 {
|
||||
return nil, errors.New("model must be provided once")
|
||||
}
|
||||
|
||||
values := gjson.GetManyBytes(requestBody, "model", "group")
|
||||
model, err := getJSONStringValue(values[0], "model")
|
||||
@@ -240,6 +333,64 @@ func getModelFromJSONBody(c *gin.Context) (*ModelRequest, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
func countTopLevelJSONKey(data []byte, target string) int {
|
||||
depth := 0
|
||||
inString := false
|
||||
escaped := false
|
||||
stringStart := 0
|
||||
expectingKey := false
|
||||
count := 0
|
||||
for index, current := range data {
|
||||
if inString {
|
||||
if escaped {
|
||||
escaped = false
|
||||
continue
|
||||
}
|
||||
if current == '\\' {
|
||||
escaped = true
|
||||
continue
|
||||
}
|
||||
if current != '"' {
|
||||
continue
|
||||
}
|
||||
inString = false
|
||||
if depth == 1 && expectingKey {
|
||||
key := string(data[stringStart:index])
|
||||
var decodedKey string
|
||||
if common.Unmarshal(data[stringStart-1:index+1], &decodedKey) == nil {
|
||||
key = decodedKey
|
||||
}
|
||||
cursor := index + 1
|
||||
for cursor < len(data) && (data[cursor] == ' ' || data[cursor] == '\t' || data[cursor] == '\r' || data[cursor] == '\n') {
|
||||
cursor++
|
||||
}
|
||||
if cursor < len(data) && data[cursor] == ':' && key == target {
|
||||
count++
|
||||
}
|
||||
expectingKey = false
|
||||
}
|
||||
continue
|
||||
}
|
||||
switch current {
|
||||
case '"':
|
||||
inString = true
|
||||
stringStart = index + 1
|
||||
case '{':
|
||||
depth++
|
||||
if depth == 1 {
|
||||
expectingKey = true
|
||||
}
|
||||
case '}':
|
||||
depth--
|
||||
case ',':
|
||||
if depth == 1 {
|
||||
expectingKey = true
|
||||
}
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func getJSONStringValue(result gjson.Result, field string) (string, error) {
|
||||
if !result.Exists() || result.Type == gjson.Null {
|
||||
return "", nil
|
||||
@@ -254,7 +405,9 @@ func getModelRequest(c *gin.Context) (*ModelRequest, bool, error) {
|
||||
var modelRequest ModelRequest
|
||||
shouldSelectChannel := true
|
||||
var err error
|
||||
if strings.Contains(c.Request.URL.Path, "/mj/") {
|
||||
if modelName := c.GetString("resolved_task_model"); modelName != "" {
|
||||
modelRequest.Model = modelName
|
||||
} else if strings.Contains(c.Request.URL.Path, "/mj/") {
|
||||
relayMode := relayconstant.Path2RelayModeMidjourney(c.Request.URL.Path)
|
||||
if relayMode == relayconstant.RelayModeMidjourneyTaskFetch ||
|
||||
relayMode == relayconstant.RelayModeMidjourneyTaskFetchByCondition ||
|
||||
@@ -282,17 +435,6 @@ func getModelRequest(c *gin.Context) (*ModelRequest, bool, error) {
|
||||
modelRequest.Model = midjourneyModel
|
||||
}
|
||||
c.Set("relay_mode", relayMode)
|
||||
} else if strings.Contains(c.Request.URL.Path, "/suno/") {
|
||||
relayMode := relayconstant.Path2RelaySuno(c.Request.Method, c.Request.URL.Path)
|
||||
if relayMode == relayconstant.RelayModeSunoFetch ||
|
||||
relayMode == relayconstant.RelayModeSunoFetchByID {
|
||||
shouldSelectChannel = false
|
||||
} else {
|
||||
modelName := service.CoverTaskActionToModelName(constant.TaskPlatformSuno, c.Param("action"))
|
||||
modelRequest.Model = modelName
|
||||
}
|
||||
c.Set("platform", string(constant.TaskPlatformSuno))
|
||||
c.Set("relay_mode", relayMode)
|
||||
} else if strings.Contains(c.Request.URL.Path, "/v1/videos/") && strings.HasSuffix(c.Request.URL.Path, "/remix") {
|
||||
relayMode := relayconstant.RelayModeVideoSubmit
|
||||
c.Set("relay_mode", relayMode)
|
||||
@@ -423,10 +565,6 @@ func getTaskOriginModelName(c *gin.Context) string {
|
||||
}
|
||||
|
||||
taskId := c.Param("task_id")
|
||||
if taskId == "" {
|
||||
// jimeng adapter
|
||||
taskId = c.GetString("task_id")
|
||||
}
|
||||
if taskId == "" {
|
||||
return ""
|
||||
}
|
||||
@@ -440,15 +578,54 @@ func getTaskOriginModelName(c *gin.Context) string {
|
||||
|
||||
func SetupContextForSelectedChannel(c *gin.Context, channel *model.Channel, modelName string) *types.NewAPIError {
|
||||
c.Set("original_model", modelName) // for retry
|
||||
expectedPlugin := c.GetString("expected_task_plugin_key")
|
||||
if channel == nil {
|
||||
logTaskPluginChannelDecision(c, nil, modelName, "channel_rejected", "nil_channel")
|
||||
return types.NewError(errors.New("channel is nil"), types.ErrorCodeGetChannelFailed, types.ErrOptionWithSkipRetry())
|
||||
}
|
||||
if expectedPlugin != "" && !channelMatchesExpectedTaskPlugin(c, channel, expectedPlugin) {
|
||||
logTaskPluginChannelDecision(c, channel, modelName, "channel_rejected", "identity_mismatch")
|
||||
return types.NewError(
|
||||
errors.New("selected channel does not match the pinned task plugin"),
|
||||
types.ErrorCodeGetChannelFailed,
|
||||
types.ErrOptionWithSkipRetry(),
|
||||
)
|
||||
}
|
||||
if candidate, matched := pinnedEndpointCandidateForChannel(c, channel, expectedPlugin); matched {
|
||||
if value, exists := c.Get(jsplugin.ContextKeyPinnedEndpoint); exists {
|
||||
if pinned, ok := value.(jsplugin.PinnedEndpoint); ok && candidate.Plugin != nil && candidate.Plugin != pinned.Plugin {
|
||||
previousPlugin := pinned.Plugin.Meta.Key
|
||||
pinned.Plugin = candidate.Plugin
|
||||
pinned.Protocol = candidate.Protocol
|
||||
pinned.Operation = candidate.Operation
|
||||
c.Set(jsplugin.ContextKeyPinnedEndpoint, pinned)
|
||||
c.Set(jsplugin.ContextKeyPinnedPlugin, jsplugin.PinnedPlugin{Generation: pinned.Generation, Plugin: candidate.Plugin})
|
||||
c.Set("expected_task_plugin_key", candidate.Plugin.Meta.Key)
|
||||
c.Set("task_plugin_key", candidate.Plugin.Meta.Key)
|
||||
c.Set("platform", candidate.Plugin.Meta.Key)
|
||||
logger.LogDebug(
|
||||
c,
|
||||
"task_plugin subsystem=endpoint event=provider_selected generation=%d previous_plugin=%q plugin=%q model=%q channel_id=%d channel_type=%d",
|
||||
pinned.Generation.Number,
|
||||
previousPlugin,
|
||||
candidate.Plugin.Meta.Key,
|
||||
modelName,
|
||||
channel.Id,
|
||||
channel.Type,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
common.SetContextKey(c, constant.ContextKeyChannelId, channel.Id)
|
||||
common.SetContextKey(c, constant.ContextKeyChannelName, channel.Name)
|
||||
common.SetContextKey(c, constant.ContextKeyChannelType, channel.Type)
|
||||
common.SetContextKey(c, constant.ContextKeyChannelCreateTime, channel.CreatedTime)
|
||||
common.SetContextKey(c, constant.ContextKeyChannelSetting, channel.GetSetting())
|
||||
common.SetContextKey(c, constant.ContextKeyChannelOtherSetting, channel.GetOtherSettings())
|
||||
if channel.Type == constant.ChannelTypeTaskPlugin {
|
||||
c.Set("task_plugin_key", channel.GetSetting().TaskPluginKey)
|
||||
}
|
||||
logTaskPluginChannelDecision(c, channel, modelName, "channel_selected", "")
|
||||
paramOverride := channel.GetParamOverride()
|
||||
headerOverride := channel.GetHeaderOverride()
|
||||
if mergedParam, applied := service.ApplyChannelAffinityOverrideTemplate(c, paramOverride); applied {
|
||||
|
||||
Reference in New Issue
Block a user