Files
new-api/controller/plugin_native_e2e_test.go
T

245 lines
9.1 KiB
Go

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)
}