mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-09 03:28:15 +00:00
feat(task): replace built-in task adaptors with a sandboxed JS plugin system (#7076)
This commit is contained in:
@@ -0,0 +1,244 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/QuantumNous/new-api/common"
|
||||
"github.com/QuantumNous/new-api/constant"
|
||||
"github.com/QuantumNous/new-api/middleware"
|
||||
"github.com/QuantumNous/new-api/model"
|
||||
pluginruntime "github.com/QuantumNous/new-api/pkg/jsplugin"
|
||||
"github.com/QuantumNous/new-api/relay"
|
||||
relaycommon "github.com/QuantumNous/new-api/relay/common"
|
||||
"github.com/QuantumNous/new-api/service"
|
||||
"github.com/QuantumNous/new-api/setting/ratio_setting"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type nativeRouteBilling struct {
|
||||
events []string
|
||||
preConsumed int
|
||||
userID int
|
||||
settled bool
|
||||
}
|
||||
|
||||
func (b *nativeRouteBilling) Settle(int) error {
|
||||
b.events = append(b.events, "settle")
|
||||
b.settled = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *nativeRouteBilling) Refund(*gin.Context) {
|
||||
b.events = append(b.events, "refund")
|
||||
if !b.settled && b.preConsumed > 0 {
|
||||
_ = model.IncreaseUserQuota(b.userID, b.preConsumed, true)
|
||||
b.preConsumed = 0
|
||||
}
|
||||
}
|
||||
|
||||
func (b *nativeRouteBilling) NeedsRefund() bool {
|
||||
return !b.settled && b.preConsumed > 0
|
||||
}
|
||||
|
||||
func (b *nativeRouteBilling) GetPreConsumedQuota() int {
|
||||
return b.preConsumed
|
||||
}
|
||||
|
||||
func (b *nativeRouteBilling) Reserve(quota int) error {
|
||||
b.events = append(b.events, "reserve")
|
||||
if err := model.DecreaseUserQuota(b.userID, quota, true); err != nil {
|
||||
return err
|
||||
}
|
||||
b.preConsumed = quota
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestKlingNativeRouteSubmitPollSettleAndQuery(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
service.InitHttpClient()
|
||||
|
||||
previousDB := model.DB
|
||||
previousLogDB := model.LOG_DB
|
||||
previousMemoryCache := common.MemoryCacheEnabled
|
||||
previousBatchUpdate := common.BatchUpdateEnabled
|
||||
previousLogConsume := common.LogConsumeEnabled
|
||||
previousRedisEnabled := common.RedisEnabled
|
||||
database, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, database.AutoMigrate(&model.User{}, &model.Channel{}, &model.Task{}, &model.Log{}))
|
||||
model.DB = database
|
||||
model.LOG_DB = database
|
||||
common.MemoryCacheEnabled = false
|
||||
common.BatchUpdateEnabled = false
|
||||
common.LogConsumeEnabled = false
|
||||
common.RedisEnabled = false
|
||||
previousModelRatios := ratio_setting.ModelRatio2JSONString()
|
||||
require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(`{"kling-v1":1}`))
|
||||
t.Cleanup(func() {
|
||||
model.DB = previousDB
|
||||
model.LOG_DB = previousLogDB
|
||||
common.MemoryCacheEnabled = previousMemoryCache
|
||||
common.BatchUpdateEnabled = previousBatchUpdate
|
||||
common.LogConsumeEnabled = previousLogConsume
|
||||
common.RedisEnabled = previousRedisEnabled
|
||||
require.NoError(t, ratio_setting.UpdateModelRatioByJSONString(previousModelRatios))
|
||||
})
|
||||
require.NoError(t, database.Create(&model.User{
|
||||
Id: 7,
|
||||
Username: "native-route-user",
|
||||
Group: "default",
|
||||
Quota: 1_000_000,
|
||||
}).Error)
|
||||
|
||||
var submitCalls atomic.Int32
|
||||
var queryCalls atomic.Int32
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
switch {
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/kling/v1/videos/text2video":
|
||||
submitCalls.Add(1)
|
||||
body, readErr := io.ReadAll(r.Body)
|
||||
if !assert.NoError(t, readErr) {
|
||||
http.Error(w, "read request", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
assert.Contains(t, string(body), `"model_name":"kling-v1"`)
|
||||
_, _ = io.WriteString(w, `{"code":0,"message":"","data":{"task_id":"kling-private-1","task_status":"submitted"}}`)
|
||||
case r.Method == http.MethodGet && r.URL.Path == "/kling/v1/videos/text2video/kling-private-1":
|
||||
queryCalls.Add(1)
|
||||
_, _ = io.WriteString(w, `{"code":0,"message":"","data":{"task_id":"kling-private-1","task_status":"succeed","task_status_msg":"","task_result":{"videos":[{"id":"video-private","url":"https://cdn.example/video.mp4","duration":"5"}]},"final_unit_deduction":"1"}}`)
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
channel := model.Channel{
|
||||
Type: constant.ChannelTypeKling,
|
||||
Name: "kling-native-e2e",
|
||||
Key: "sk-test",
|
||||
BaseURL: &upstream.URL,
|
||||
Status: common.ChannelStatusEnabled,
|
||||
Models: "kling-v1",
|
||||
Group: "default",
|
||||
}
|
||||
require.NoError(t, database.Create(&channel).Error)
|
||||
|
||||
generation := pluginruntime.DefaultRegistry.Generation()
|
||||
require.NotNil(t, generation)
|
||||
submitBinding, found := generation.LookupDeclaredRoute(http.MethodPost, "/kling/v1/videos/text2video")
|
||||
require.True(t, found)
|
||||
require.Equal(t, "kling", submitBinding.Plugin.Meta.Key)
|
||||
|
||||
submitRecorder := httptest.NewRecorder()
|
||||
submitContext, _ := gin.CreateTestContext(submitRecorder)
|
||||
submitContext.Request = httptest.NewRequest(
|
||||
http.MethodPost,
|
||||
"/kling/v1/videos/text2video",
|
||||
bytes.NewBufferString(`{"model_name":"kling-v1","prompt":"a lighthouse"}`),
|
||||
)
|
||||
submitContext.Request.Header.Set("Content-Type", "application/json")
|
||||
submitContext.Set(pluginruntime.ContextKeyPinnedRoute, pluginruntime.PinnedRoute{
|
||||
Generation: generation,
|
||||
Plugin: submitBinding.Plugin,
|
||||
Route: submitBinding.Route,
|
||||
})
|
||||
common.SetContextKey(submitContext, constant.ContextKeyUserId, 7)
|
||||
common.SetContextKey(submitContext, constant.ContextKeyUserGroup, "default")
|
||||
common.SetContextKey(submitContext, constant.ContextKeyUsingGroup, "default")
|
||||
common.SetContextKey(submitContext, constant.ContextKeyTokenGroup, "default")
|
||||
common.SetContextKey(submitContext, constant.ContextKeyUserQuota, 1_000_000)
|
||||
|
||||
middleware.PrepareTaskPluginRoute()(submitContext)
|
||||
require.False(t, submitContext.IsAborted(), submitRecorder.Body.String())
|
||||
require.Equal(t, "kling-v1", submitContext.GetString("resolved_task_model"))
|
||||
require.Equal(t, "text_to_video", submitContext.GetString("task_action"))
|
||||
require.Nil(t, middleware.SetupContextForSelectedChannel(submitContext, &channel, "kling-v1"))
|
||||
|
||||
billing := &nativeRouteBilling{userID: 7}
|
||||
relayInfo := &relaycommon.RelayInfo{
|
||||
UserId: 7,
|
||||
UserGroup: "default",
|
||||
UsingGroup: "default",
|
||||
UserQuota: 1_000_000,
|
||||
TokenGroup: "default",
|
||||
OriginModelName: "kling-v1",
|
||||
Billing: billing,
|
||||
TaskRelayInfo: &relaycommon.TaskRelayInfo{
|
||||
Action: submitContext.GetString("task_action"),
|
||||
PublicTaskID: "task_kling_public",
|
||||
LockedChannel: &channel,
|
||||
},
|
||||
}
|
||||
|
||||
outcome, taskErr := executeTaskSubmissionWith(submitContext, relayInfo, relay.RelayTaskSubmit)
|
||||
require.Nil(t, taskErr)
|
||||
require.NotNil(t, outcome)
|
||||
require.Equal(t, []string{"reserve", "settle"}, billing.events)
|
||||
require.False(t, submitContext.Writer.Written())
|
||||
|
||||
presentTaskSubmission(submitContext, outcome)
|
||||
require.Equal(t, http.StatusOK, submitRecorder.Code)
|
||||
assert.Contains(t, submitRecorder.Body.String(), `"task_id":"task_kling_public"`)
|
||||
assert.NotContains(t, submitRecorder.Body.String(), "kling-private-1")
|
||||
assert.Equal(t, int32(1), submitCalls.Load())
|
||||
|
||||
var persisted model.Task
|
||||
require.NoError(t, database.Where("task_id = ?", "task_kling_public").First(&persisted).Error)
|
||||
assert.Equal(t, constant.TaskPlatform("kling"), persisted.Platform)
|
||||
assert.Equal(t, "kling-private-1", persisted.PrivateData.UpstreamTaskID)
|
||||
assert.Equal(t, model.TaskStatus(model.TaskStatusNotStart), persisted.Status)
|
||||
|
||||
previousAdaptorFactory := service.GetTaskAdaptorFunc
|
||||
service.GetTaskAdaptorFunc = func(platform constant.TaskPlatform) service.TaskPollingAdaptor {
|
||||
return relay.GetTaskAdaptor(platform)
|
||||
}
|
||||
t.Cleanup(func() { service.GetTaskAdaptorFunc = previousAdaptorFactory })
|
||||
service.DispatchPlatformUpdate(
|
||||
context.Background(),
|
||||
persisted.Platform,
|
||||
map[int][]string{channel.Id: {"kling-private-1"}},
|
||||
map[string]*model.Task{"kling-private-1": &persisted},
|
||||
)
|
||||
|
||||
require.NoError(t, database.Where("task_id = ?", "task_kling_public").First(&persisted).Error)
|
||||
assert.Equal(t, model.TaskStatus(model.TaskStatusSuccess), persisted.Status)
|
||||
assert.Equal(t, "100%", persisted.Progress)
|
||||
assert.Equal(t, 1, persisted.Quota)
|
||||
assert.Equal(t, int32(1), queryCalls.Load())
|
||||
var settledUser model.User
|
||||
require.NoError(t, database.First(&settledUser, 7).Error)
|
||||
assert.Equal(t, 999_999, settledUser.Quota)
|
||||
|
||||
queryBinding, found := generation.LookupDeclaredRoute(http.MethodGet, "/kling/v1/videos/text2video/:task_id")
|
||||
require.True(t, found)
|
||||
queryRecorder := httptest.NewRecorder()
|
||||
queryContext, _ := gin.CreateTestContext(queryRecorder)
|
||||
queryContext.Request = httptest.NewRequest(http.MethodGet, "/kling/v1/videos/text2video/task_kling_public", nil)
|
||||
queryContext.Params = gin.Params{{Key: "task_id", Value: "task_kling_public"}}
|
||||
queryContext.Set(pluginruntime.ContextKeyPinnedRoute, pluginruntime.PinnedRoute{
|
||||
Generation: generation,
|
||||
Plugin: queryBinding.Plugin,
|
||||
Route: queryBinding.Route,
|
||||
})
|
||||
common.SetContextKey(queryContext, constant.ContextKeyUserId, 7)
|
||||
|
||||
middleware.PrepareTaskPluginRoute()(queryContext)
|
||||
|
||||
require.True(t, queryContext.IsAborted())
|
||||
require.Equal(t, http.StatusOK, queryRecorder.Code)
|
||||
assert.Contains(t, queryRecorder.Body.String(), `"task_id":"task_kling_public"`)
|
||||
assert.Contains(t, queryRecorder.Body.String(), `"task_status":"succeed"`)
|
||||
assert.NotContains(t, queryRecorder.Body.String(), "kling-private-1")
|
||||
assert.NotContains(t, queryRecorder.Body.String(), upstream.URL)
|
||||
}
|
||||
Reference in New Issue
Block a user