mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-08-31 02:41:34 +00:00
390 lines
15 KiB
Go
390 lines
15 KiB
Go
package controller
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
"github.com/QuantumNous/new-api/constant"
|
|
"github.com/QuantumNous/new-api/dto"
|
|
"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/types"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/glebarez/sqlite"
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
type taskSubmissionTestBilling struct {
|
|
events *[]string
|
|
settleErr error
|
|
onSettle func()
|
|
refunds int
|
|
}
|
|
|
|
func (b *taskSubmissionTestBilling) Settle(int) error {
|
|
*b.events = append(*b.events, "settle")
|
|
if b.onSettle != nil {
|
|
b.onSettle()
|
|
}
|
|
return b.settleErr
|
|
}
|
|
|
|
func (b *taskSubmissionTestBilling) Refund(*gin.Context) {
|
|
*b.events = append(*b.events, "refund")
|
|
b.refunds++
|
|
}
|
|
|
|
func (b *taskSubmissionTestBilling) NeedsRefund() bool { return b.refunds == 0 }
|
|
func (b *taskSubmissionTestBilling) GetPreConsumedQuota() int { return 0 }
|
|
func (b *taskSubmissionTestBilling) Reserve(int) error {
|
|
*b.events = append(*b.events, "reserve")
|
|
return nil
|
|
}
|
|
|
|
func TestPresentTaskSubmissionUsesNativePresenterAfterPersistence(t *testing.T) {
|
|
plugin, err := pluginruntime.CompilePlugin(`
|
|
export const meta = {apiVersion:1,key:"presenter-test",name:"Presenter",version:"1.0.0",author:{name:"Test"},models:["model"],fetchMode:"per_task",routes:[{method:"POST",path:"/vendor/jobs",type:"submit",decode:"decode",render:"created"}]};
|
|
export const native = {decode:function(ctx){return {kind:"submit",model:"model",requestBody:ctx.body.value};},created:function(ctx,task){return {data:{task_id:task.task_id},upstream:task.data};}};
|
|
export function buildSubmitRequest(){return {}} export function parseSubmitResponse(){return {taskId:"upstream"}} export function buildQueryRequest(){return {}} export function parseTaskResult(){return {status:"SUCCESS"}}
|
|
`, pluginruntime.Options{})
|
|
require.NoError(t, err)
|
|
priceData := types.PriceData{}
|
|
priceData.AddOtherRatio("seconds", 5)
|
|
task := &model.Task{TaskID: "task_public", SubmitTime: 123}
|
|
task.SetData(map[string]any{"task_id": "upstream_private"})
|
|
outcome := &taskSubmissionOutcome{
|
|
Result: &relay.TaskSubmitResult{},
|
|
Task: task,
|
|
RelayInfo: &relaycommon.RelayInfo{PriceData: priceData},
|
|
}
|
|
recorder := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(recorder)
|
|
c.Request = httptest.NewRequest(http.MethodPost, "/vendor/jobs", strings.NewReader(`{"model":"model"}`))
|
|
c.Set(pluginruntime.ContextKeyPinnedRoute, pluginruntime.PinnedRoute{Plugin: plugin, Route: plugin.Meta.Routes[0]})
|
|
c.Set(pluginruntime.ContextKeyRouteRequest, pluginruntime.RouteRequestContext{Path: "/vendor/jobs", Method: http.MethodPost, Body: map[string]any{"kind": "json", "value": map[string]any{"model": "model"}}})
|
|
|
|
presentTaskSubmission(c, outcome)
|
|
|
|
assert.JSONEq(t, `{
|
|
"data":{"task_id":"task_public"},
|
|
"upstream":{"task_id":"upstream_private"}
|
|
}`, recorder.Body.String())
|
|
assert.JSONEq(t, `{"seconds":5}`, recorder.Header().Get("X-New-Api-Other-Ratios"))
|
|
}
|
|
|
|
func TestPresentTaskSubmissionFallbackUsesPersistedPublicID(t *testing.T) {
|
|
recorder := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(recorder)
|
|
outcome := &taskSubmissionOutcome{
|
|
Result: &relay.TaskSubmitResult{},
|
|
Task: &model.Task{TaskID: "task_persisted", SubmitTime: 456},
|
|
RelayInfo: &relaycommon.RelayInfo{OriginModelName: "video-model"},
|
|
}
|
|
|
|
presentTaskSubmission(c, outcome)
|
|
|
|
assert.JSONEq(t, `{
|
|
"id":"task_persisted",
|
|
"task_id":"task_persisted",
|
|
"status":"queued",
|
|
"model":"video-model",
|
|
"created_at":456
|
|
}`, recorder.Body.String())
|
|
}
|
|
|
|
func TestPresentTaskSubmissionUsesHostOpenAIVideoCreateReceipt(t *testing.T) {
|
|
recorder := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(recorder)
|
|
c.Set(pluginruntime.ContextKeyPinnedEndpoint, pluginruntime.PinnedEndpoint{
|
|
Protocol: "openai_video",
|
|
Operation: pluginruntime.HostProtocolOperation{Name: "create"},
|
|
})
|
|
task := &model.Task{
|
|
TaskID: "task_public",
|
|
Status: model.TaskStatusSubmitted,
|
|
Progress: "0%",
|
|
CreatedAt: 456,
|
|
Properties: model.Properties{OriginModelName: "video-model"},
|
|
}
|
|
outcome := &taskSubmissionOutcome{Result: &relay.TaskSubmitResult{}, Task: task, RelayInfo: &relaycommon.RelayInfo{}}
|
|
|
|
presentTaskSubmission(c, outcome)
|
|
|
|
assert.JSONEq(t, `{"id":"task_public","object":"video","model":"video-model","status":"queued","progress":0,"created_at":456}`, recorder.Body.String())
|
|
assert.NotContains(t, recorder.Body.String(), "task_id")
|
|
}
|
|
|
|
func TestExecuteTaskSubmissionRefundsWhenInsertFails(t *testing.T) {
|
|
events := make([]string, 0, 3)
|
|
database := setupTaskSubmissionDatabase(t, false, &events)
|
|
_ = database
|
|
billing := &taskSubmissionTestBilling{events: &events}
|
|
c := taskSubmissionTestContext()
|
|
info := taskSubmissionRelayInfo(billing)
|
|
|
|
outcome, taskErr := executeTaskSubmissionWith(c, info, func(*gin.Context, *relaycommon.RelayInfo) (*relay.TaskSubmitResult, *dto.TaskError) {
|
|
return &relay.TaskSubmitResult{
|
|
UpstreamTaskID: "upstream_private",
|
|
Platform: constant.TaskPlatform("plugin"),
|
|
}, nil
|
|
})
|
|
|
|
assert.Nil(t, outcome)
|
|
require.NotNil(t, taskErr)
|
|
assert.Equal(t, "task_insert_failed", taskErr.Code)
|
|
assert.Equal(t, []string{"reserve", "insert", "refund"}, events)
|
|
assert.Equal(t, 1, billing.refunds)
|
|
assert.False(t, c.Writer.Written())
|
|
}
|
|
|
|
func TestExecuteTaskSubmissionSettlementFailureStaysDurableAndWritesNothing(t *testing.T) {
|
|
events := make([]string, 0, 3)
|
|
database := setupTaskSubmissionDatabase(t, true, &events)
|
|
billing := &taskSubmissionTestBilling{events: &events, settleErr: errors.New("settlement failed")}
|
|
c := taskSubmissionTestContext()
|
|
info := taskSubmissionRelayInfo(billing)
|
|
|
|
outcome, taskErr := executeTaskSubmissionWith(c, info, func(*gin.Context, *relaycommon.RelayInfo) (*relay.TaskSubmitResult, *dto.TaskError) {
|
|
return &relay.TaskSubmitResult{
|
|
UpstreamTaskID: "upstream_private",
|
|
Platform: constant.TaskPlatform("plugin"),
|
|
}, nil
|
|
})
|
|
|
|
assert.Nil(t, outcome)
|
|
require.NotNil(t, taskErr)
|
|
assert.Equal(t, "task_billing_settlement_failed", taskErr.Code)
|
|
assert.Equal(t, []string{"reserve", "insert", "settle"}, events)
|
|
assert.Zero(t, billing.refunds)
|
|
var count int64
|
|
require.NoError(t, database.Model(&model.Task{}).Where("task_id = ?", "task_public").Count(&count).Error)
|
|
assert.Equal(t, int64(1), count)
|
|
assert.False(t, c.Writer.Written())
|
|
}
|
|
|
|
func TestExecuteTaskSubmissionPersistsPinnedPluginProvenance(t *testing.T) {
|
|
events := make([]string, 0, 3)
|
|
database := setupTaskSubmissionDatabase(t, true, &events)
|
|
previousLogConsumeEnabled := common.LogConsumeEnabled
|
|
common.LogConsumeEnabled = false
|
|
t.Cleanup(func() { common.LogConsumeEnabled = previousLogConsumeEnabled })
|
|
|
|
c := taskSubmissionTestContext()
|
|
c.Set(common.RequestIdKey, "request-public")
|
|
c.Set(pluginruntime.ContextKeyPinnedPlugin, pluginruntime.PinnedPlugin{
|
|
Generation: &pluginruntime.RoutingGeneration{Number: 42},
|
|
Plugin: &pluginruntime.LoadedPlugin{Meta: pluginruntime.Meta{
|
|
Key: "document-parser",
|
|
Name: "Document Parser",
|
|
Version: "1.2.3",
|
|
APIVersion: 1,
|
|
Author: pluginruntime.AuthorMeta{
|
|
Name: "Community Author",
|
|
URL: "https://plugins.example/author",
|
|
},
|
|
}},
|
|
})
|
|
billing := &taskSubmissionTestBilling{events: &events}
|
|
info := taskSubmissionRelayInfo(billing)
|
|
|
|
outcome, taskErr := executeTaskSubmissionWith(c, info, func(*gin.Context, *relaycommon.RelayInfo) (*relay.TaskSubmitResult, *dto.TaskError) {
|
|
return &relay.TaskSubmitResult{
|
|
UpstreamTaskID: "upstream-private",
|
|
Platform: constant.TaskPlatform("document-parser"),
|
|
}, nil
|
|
})
|
|
|
|
require.Nil(t, taskErr)
|
|
require.NotNil(t, outcome)
|
|
require.NotNil(t, outcome.Task.PrivateData.Execution)
|
|
require.NotNil(t, outcome.Task.PrivateData.Execution.TaskPlugin)
|
|
assert.Equal(t, "request-public", outcome.Task.PrivateData.Execution.RequestID)
|
|
assert.Equal(t, "/plugin/submit", outcome.Task.PrivateData.Execution.RequestPath)
|
|
assert.Equal(t, "1.2.3", outcome.Task.PrivateData.Execution.TaskPlugin.Version)
|
|
assert.Equal(t, uint64(42), outcome.Task.PrivateData.Execution.TaskPlugin.Generation)
|
|
require.NotNil(t, outcome.Task.PrivateData.Execution.TaskPlugin.Author)
|
|
assert.Equal(t, "Community Author", outcome.Task.PrivateData.Execution.TaskPlugin.Author.Name)
|
|
assert.Equal(t, "https://plugins.example/author", outcome.Task.PrivateData.Execution.TaskPlugin.Author.URL)
|
|
|
|
var stored model.Task
|
|
require.NoError(t, database.Where("task_id = ?", "task_public").First(&stored).Error)
|
|
require.NotNil(t, stored.PrivateData.Execution)
|
|
require.NotNil(t, stored.PrivateData.Execution.TaskPlugin)
|
|
assert.Equal(t, "document-parser", stored.PrivateData.Execution.TaskPlugin.Key)
|
|
require.NotNil(t, stored.PrivateData.Execution.TaskPlugin.Author)
|
|
assert.Equal(t, "Community Author", stored.PrivateData.Execution.TaskPlugin.Author.Name)
|
|
assert.Equal(t, "upstream-private", stored.PrivateData.UpstreamTaskID)
|
|
}
|
|
|
|
func TestExecuteTaskSubmissionRefundsCancellationBeforeDurableBarrier(t *testing.T) {
|
|
events := make([]string, 0, 2)
|
|
setupTaskSubmissionDatabase(t, true, &events)
|
|
billing := &taskSubmissionTestBilling{events: &events}
|
|
c := taskSubmissionTestContext()
|
|
requestContext, cancel := context.WithCancel(c.Request.Context())
|
|
c.Request = c.Request.WithContext(requestContext)
|
|
info := taskSubmissionRelayInfo(billing)
|
|
|
|
outcome, taskErr := executeTaskSubmissionWith(c, info, func(*gin.Context, *relaycommon.RelayInfo) (*relay.TaskSubmitResult, *dto.TaskError) {
|
|
cancel()
|
|
return &relay.TaskSubmitResult{
|
|
UpstreamTaskID: "upstream_private",
|
|
Platform: constant.TaskPlatform("plugin"),
|
|
}, nil
|
|
})
|
|
|
|
assert.Nil(t, outcome)
|
|
require.NotNil(t, taskErr)
|
|
assert.Equal(t, "request_cancelled", taskErr.Code)
|
|
assert.Equal(t, []string{"refund"}, events)
|
|
assert.Equal(t, 1, billing.refunds)
|
|
assert.False(t, c.Writer.Written())
|
|
}
|
|
|
|
func TestExecuteTaskSubmissionDisconnectBeforeUpstreamAcceptanceSkipsSubmitAndRefunds(t *testing.T) {
|
|
events := make([]string, 0, 1)
|
|
setupTaskSubmissionDatabase(t, true, &events)
|
|
billing := &taskSubmissionTestBilling{events: &events}
|
|
c := taskSubmissionTestContext()
|
|
requestContext, cancel := context.WithCancel(c.Request.Context())
|
|
cancel()
|
|
c.Request = c.Request.WithContext(requestContext)
|
|
info := taskSubmissionRelayInfo(billing)
|
|
submitted := false
|
|
|
|
outcome, taskErr := executeTaskSubmissionWith(c, info, func(*gin.Context, *relaycommon.RelayInfo) (*relay.TaskSubmitResult, *dto.TaskError) {
|
|
submitted = true
|
|
return nil, nil
|
|
})
|
|
|
|
assert.Nil(t, outcome)
|
|
require.NotNil(t, taskErr)
|
|
assert.Equal(t, "request_cancelled", taskErr.Code)
|
|
assert.False(t, submitted)
|
|
assert.Equal(t, []string{"refund"}, events)
|
|
assert.Equal(t, 1, billing.refunds)
|
|
assert.False(t, c.Writer.Written())
|
|
}
|
|
|
|
func TestExecuteTaskSubmissionCallerCancellationDuringSubmitRefundsBeforeDurableBarrier(t *testing.T) {
|
|
events := make([]string, 0, 1)
|
|
setupTaskSubmissionDatabase(t, true, &events)
|
|
billing := &taskSubmissionTestBilling{events: &events}
|
|
c := taskSubmissionTestContext()
|
|
requestContext, cancel := context.WithCancel(c.Request.Context())
|
|
c.Request = c.Request.WithContext(requestContext)
|
|
info := taskSubmissionRelayInfo(billing)
|
|
submitStarted := make(chan struct{})
|
|
done := make(chan struct{})
|
|
var outcome *taskSubmissionOutcome
|
|
var taskErr *dto.TaskError
|
|
|
|
go func() {
|
|
defer close(done)
|
|
outcome, taskErr = executeTaskSubmissionWith(c, info, func(c *gin.Context, _ *relaycommon.RelayInfo) (*relay.TaskSubmitResult, *dto.TaskError) {
|
|
close(submitStarted)
|
|
<-c.Request.Context().Done()
|
|
return nil, service.TaskErrorWrapperLocal(c.Request.Context().Err(), "do_request_failed", http.StatusInternalServerError)
|
|
})
|
|
}()
|
|
select {
|
|
case <-submitStarted:
|
|
case <-time.After(2 * time.Second):
|
|
require.FailNow(t, "submission did not start")
|
|
}
|
|
cancel()
|
|
select {
|
|
case <-done:
|
|
case <-time.After(2 * time.Second):
|
|
require.FailNow(t, "submission did not stop after disconnect")
|
|
}
|
|
|
|
assert.Nil(t, outcome)
|
|
require.NotNil(t, taskErr)
|
|
assert.Equal(t, "request_cancelled", taskErr.Code)
|
|
assert.Equal(t, []string{"refund"}, events)
|
|
assert.Equal(t, 1, billing.refunds)
|
|
assert.False(t, c.Writer.Written())
|
|
}
|
|
|
|
func TestExecuteTaskSubmissionDisconnectAfterDurableInsertDoesNotRefund(t *testing.T) {
|
|
events := make([]string, 0, 3)
|
|
database := setupTaskSubmissionDatabase(t, true, &events)
|
|
previousLogConsumeEnabled := common.LogConsumeEnabled
|
|
common.LogConsumeEnabled = false
|
|
t.Cleanup(func() { common.LogConsumeEnabled = previousLogConsumeEnabled })
|
|
c := taskSubmissionTestContext()
|
|
requestContext, cancel := context.WithCancel(c.Request.Context())
|
|
c.Request = c.Request.WithContext(requestContext)
|
|
billing := &taskSubmissionTestBilling{
|
|
events: &events,
|
|
onSettle: cancel,
|
|
}
|
|
info := taskSubmissionRelayInfo(billing)
|
|
|
|
outcome, taskErr := executeTaskSubmissionWith(c, info, func(*gin.Context, *relaycommon.RelayInfo) (*relay.TaskSubmitResult, *dto.TaskError) {
|
|
return &relay.TaskSubmitResult{
|
|
UpstreamTaskID: "upstream_private",
|
|
Platform: constant.TaskPlatform("plugin"),
|
|
}, nil
|
|
})
|
|
|
|
require.Nil(t, taskErr)
|
|
require.NotNil(t, outcome)
|
|
assert.Equal(t, "task_public", outcome.Task.TaskID)
|
|
assert.Equal(t, []string{"reserve", "insert", "settle"}, events)
|
|
assert.Zero(t, billing.refunds)
|
|
var count int64
|
|
require.NoError(t, database.Model(&model.Task{}).Where("task_id = ?", "task_public").Count(&count).Error)
|
|
assert.Equal(t, int64(1), count)
|
|
assert.False(t, c.Writer.Written())
|
|
}
|
|
|
|
func setupTaskSubmissionDatabase(t *testing.T, migrate bool, events *[]string) *gorm.DB {
|
|
t.Helper()
|
|
previousDB := model.DB
|
|
database, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
|
require.NoError(t, err)
|
|
require.NoError(t, database.Callback().Create().Before("gorm:create").Register("test:task-submit-order", func(*gorm.DB) {
|
|
*events = append(*events, "insert")
|
|
}))
|
|
if migrate {
|
|
require.NoError(t, database.AutoMigrate(&model.Task{}))
|
|
}
|
|
model.DB = database
|
|
t.Cleanup(func() { model.DB = previousDB })
|
|
return database
|
|
}
|
|
|
|
func taskSubmissionTestContext() *gin.Context {
|
|
recorder := httptest.NewRecorder()
|
|
c, _ := gin.CreateTestContext(recorder)
|
|
c.Request = httptest.NewRequest(http.MethodPost, "/plugin/submit", strings.NewReader(`{}`))
|
|
return c
|
|
}
|
|
|
|
func taskSubmissionRelayInfo(billing relaycommon.BillingSettler) *relaycommon.RelayInfo {
|
|
return &relaycommon.RelayInfo{
|
|
UserId: 1,
|
|
UsingGroup: "default",
|
|
OriginModelName: "plugin-model",
|
|
Billing: billing,
|
|
TaskRelayInfo: &relaycommon.TaskRelayInfo{
|
|
PublicTaskID: "task_public",
|
|
LockedChannel: &model.Channel{Id: 1, Type: constant.ChannelTypeTaskPlugin, Name: "plugin"},
|
|
},
|
|
ChannelMeta: &relaycommon.ChannelMeta{ChannelId: 1, ChannelType: constant.ChannelTypeTaskPlugin},
|
|
}
|
|
}
|