mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-05 10:12:59 +00:00
245 lines
9.1 KiB
Go
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)
|
|
}
|