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