Files
new-api/service/task_polling_test.go
Calcium-Ion 86ac0f7745 refactor: extract protocol conversion layer into standalone relaykit module (#6369)
* test(relayconvert): add golden snapshot matrix and relaykit boundary guard

Phase 0 of the relaykit extraction plan: pin byte-level output of every
registered (from,to) request/response/stream conversion route, and
forbid kit-bound packages from growing host-only imports.

* wip(relayconvert): drop gin.Context from converter signatures; add convmeta draft

Phase 1 in progress: relayconvert now takes context.Context; host media
resolver adapts gin.Context back at the service boundary.

* refactor(relayconvert): decouple converters from RelayInfo, gin, and settings

Phase 1 of the relaykit extraction plan:
- converters now depend on convmeta.Meta (implemented by RelayInfo) instead
  of *relaycommon.RelayInfo; ClaudeConvertInfo and the format guesser move
  to convmeta with aliases left behind
- host settings reach converters via a convmeta.Options snapshot built in
  RelayInfo.ConvOptions; no more model_setting/reasoning global reads inside
  the conversion layer
- effort-suffix helpers move to service/relayconvert/reasoning (old package
  forwards); chat-to-responses upgrade policy moves to service (host routing
  logic, not conversion)
- golden conversion matrix unchanged

* test(relayconvert): tighten boundary — kit packages now free of gin/setting imports

* refactor(dto): drop gin and logger dependencies

Phase 2 (part 1): dto.Request.IsStream now takes *http.Request instead of
*gin.Context (Gemini's impl reads query/path off the std request); dto's
three logger calls become common.SysError. Boundary test allowlist is now
empty — kit-bound packages import no gin/setting/logger/model.

* refactor(kit): extract dependency-free kitutil; dto/types/relayconvert stop importing common

Phase 2 of the relaykit extraction plan:
- new service/relayconvert/kitutil holds the pure helpers the kit needs
  (JSON wrappers, pointer/string/uuid/timestamp utils, MaskSensitiveInfo,
  pluggable LogInfo/LogError hooks, Debug flag)
- dto, types, and all relayconvert packages now use kitutil; their only
  remaining internal deps are dto/types/constant
- common keeps every original symbol (MaskSensitiveInfo delegates to
  kitutil) so host code is untouched; main.go routes kit logging into
  common.SysLog/SysError and mirrors DebugEnabled
- golden conversion matrix unchanged

* refactor(kit): move EndpointType/FinishReason to types; OpenRouter dialect via Options

Kit packages (dto/types/relayconvert/reasonmap) no longer import constant:
- EndpointType and finish-reason values live in types; constant re-exports
- the OpenRouter special-case in claude->openai request conversion reads
  Options.OpenRouterDialect, set by the host from the channel type;
  InitChannelMeta invalidates the cached snapshot on channel switch

* refactor: extract relaykit submodule (dto/types/relayconvert/reasonmap)

Phase 3 of the relaykit extraction plan:
- new go module github.com/QuantumNous/new-api/relaykit containing dto
  (minus task family), types, relayconvert (with convmeta/kitutil/reasoning),
  and reasonmap; host consumes it via require + replace, go.work for dev
- task-family dto (task/suno/midjourney/video) stays in the host dto
  package; dual-consumer host files alias it as taskdto
- relaykit builds and tests standalone (GOWORK=off): no host imports,
  no gin, no DB, no settings
- golden conversion matrix unchanged

* build(docker): copy relaykit/go.mod before go mod download

The local-replace submodule's go.mod must exist inside the build context
for the main module graph to resolve.

* fix: address relaykit extraction regressions

* fix: address relaykit review regressions

* docs: document Meta nil receiver contract

* fix(relaykit): fail OpenAI→Claude conversion without max_tokens; reject negative default_max_tokens

The Claude Messages API requires max_tokens (omitting it is a 400
"Field required"), but with a nil Options.Claude.DefaultMaxTokens hook
the converters silently emitted a request the upstream is guaranteed to
reject. Both OpenAI Chat and Responses → Claude conversions now return
sharedclaude.ErrMissingMaxTokens when no path (client value, default
hook, thinking-adapter floor) supplied one. Unreachable in the host,
which always configures the hook.

Host side, claude.default_max_tokens now rejects negative values at the
option API before persisting — they would wrap into huge unsigned values
during conversion. Zero stays allowed: the current API treats
max_tokens: 0 as cache pre-warming.

* fix: make Gemini safety settings read path race-free
2026-07-27 15:56:21 +08:00

499 lines
16 KiB
Go

package service
import (
"bytes"
"context"
"io"
"net/http"
"sync"
"testing"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
taskdto "github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/model"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/bytedance/gopkg/util/gopool"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type taskPollingFetchAdaptor struct {
mu sync.Mutex
taskIDs []string
fetched chan string
blockTaskID string
blockStarted chan struct{}
releaseBlock chan struct{}
blockOnce sync.Once
}
type sunoFailurePollingAdaptor struct {
failReason string
}
func (a *sunoFailurePollingAdaptor) Init(_ *relaycommon.RelayInfo) {}
func (a *sunoFailurePollingAdaptor) FetchTask(_ string, _ string, body map[string]any, _ string) (*http.Response, error) {
taskIDs, _ := body["ids"].([]string)
items := make([]taskdto.SunoDataResponse, 0, len(taskIDs))
for _, taskID := range taskIDs {
items = append(items, taskdto.SunoDataResponse{
TaskID: taskID,
Status: string(model.TaskStatusFailure),
FailReason: a.failReason,
FinishTime: time.Now().Unix(),
})
}
responseBody, err := common.Marshal(taskdto.TaskResponse[[]taskdto.SunoDataResponse]{
Code: taskdto.TaskSuccessCode,
Data: items,
})
if err != nil {
return nil, err
}
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(responseBody)),
}, nil
}
func (a *sunoFailurePollingAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) {
return nil, nil
}
func (a *sunoFailurePollingAdaptor) AdjustBillingOnComplete(_ *model.Task, _ *relaycommon.TaskInfo) int {
return 0
}
func (a *taskPollingFetchAdaptor) Init(_ *relaycommon.RelayInfo) {}
func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, body map[string]any, _ string) (*http.Response, error) {
taskID, _ := body["task_id"].(string)
if taskID == a.blockTaskID && a.releaseBlock != nil {
a.blockOnce.Do(func() {
if a.blockStarted != nil {
close(a.blockStarted)
}
})
<-a.releaseBlock
}
a.mu.Lock()
a.taskIDs = append(a.taskIDs, taskID)
a.mu.Unlock()
if a.fetched != nil {
select {
case a.fetched <- taskID:
default:
}
}
response := taskdto.TaskResponse[model.Task]{
Code: taskdto.TaskSuccessCode,
Data: model.Task{
TaskID: taskID,
Status: model.TaskStatusInProgress,
Progress: "30%",
},
}
responseBody, err := common.Marshal(response)
if err != nil {
return nil, err
}
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(responseBody)),
}, nil
}
func (a *taskPollingFetchAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) {
return &relaycommon.TaskInfo{Status: model.TaskStatusInProgress}, nil
}
func (a *taskPollingFetchAdaptor) AdjustBillingOnComplete(_ *model.Task, _ *relaycommon.TaskInfo) int {
return 0
}
func (a *taskPollingFetchAdaptor) fetchCount() int {
a.mu.Lock()
defer a.mu.Unlock()
return len(a.taskIDs)
}
func (a *taskPollingFetchAdaptor) fetchedTaskIDs() []string {
a.mu.Lock()
defer a.mu.Unlock()
return append([]string(nil), a.taskIDs...)
}
func seedTaskPollingChannel(t *testing.T, id int, disableSleep bool) {
t.Helper()
ch := &model.Channel{
Id: id,
Type: constant.ChannelTypeKling,
Name: "polling_channel",
Key: "sk-test",
Status: common.ChannelStatusEnabled,
}
if disableSleep {
ch.SetOtherSettings(dto.ChannelOtherSettings{DisableTaskPollingSleep: true})
}
require.NoError(t, model.DB.Create(ch).Error)
}
func seedPollingTask(t *testing.T, channelID int, publicID string, upstreamID string) *model.Task {
t.Helper()
task := &model.Task{
TaskID: publicID,
Platform: constant.TaskPlatform("kling"),
UserId: 1,
ChannelId: channelID,
Action: constant.TaskActionGenerate,
Status: model.TaskStatusInProgress,
Progress: "30%",
CreatedAt: time.Now().Unix(),
UpdatedAt: time.Now().Unix(),
PrivateData: model.TaskPrivateData{
UpstreamTaskID: upstreamID,
},
}
require.NoError(t, model.DB.Create(task).Error)
return task
}
func TestUpdateVideoTasksDefaultSleepWaitsBetweenTasks(t *testing.T) {
truncate(t)
const channelID = 101
seedTaskPollingChannel(t, channelID, false)
first := seedPollingTask(t, channelID, "task_public_1", "upstream_1")
second := seedPollingTask(t, channelID, "task_public_2", "upstream_2")
adaptor := &taskPollingFetchAdaptor{}
previousFactory := GetTaskAdaptorFunc
GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor { return adaptor }
t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory })
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()
err := UpdateVideoTasks(ctx, constant.TaskPlatform("kling"), map[int][]string{
channelID: {
first.GetUpstreamTaskID(),
second.GetUpstreamTaskID(),
},
}, map[string]*model.Task{
first.GetUpstreamTaskID(): first,
second.GetUpstreamTaskID(): second,
})
require.ErrorIs(t, err, context.DeadlineExceeded)
assert.Equal(t, 1, adaptor.fetchCount())
}
func TestUpdateVideoTasksCanSkipPollingSleepPerChannel(t *testing.T) {
truncate(t)
const channelID = 102
seedTaskPollingChannel(t, channelID, true)
first := seedPollingTask(t, channelID, "task_public_3", "upstream_3")
second := seedPollingTask(t, channelID, "task_public_4", "upstream_4")
adaptor := &taskPollingFetchAdaptor{}
previousFactory := GetTaskAdaptorFunc
GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor { return adaptor }
t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory })
ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
defer cancel()
err := UpdateVideoTasks(ctx, constant.TaskPlatform("kling"), map[int][]string{
channelID: {
first.GetUpstreamTaskID(),
second.GetUpstreamTaskID(),
},
}, map[string]*model.Task{
first.GetUpstreamTaskID(): first,
second.GetUpstreamTaskID(): second,
})
require.NoError(t, err)
assert.Equal(t, 2, adaptor.fetchCount())
}
func TestUpdateVideoTasksDefaultSleepDoesNotBlockOtherChannels(t *testing.T) {
truncate(t)
const firstChannelID = 201
const secondChannelID = 202
seedTaskPollingChannel(t, firstChannelID, false)
seedTaskPollingChannel(t, secondChannelID, false)
firstChannelFirst := seedPollingTask(t, firstChannelID, "task_public_5", "upstream_a_1")
firstChannelSecond := seedPollingTask(t, firstChannelID, "task_public_6", "upstream_a_2")
secondChannelFirst := seedPollingTask(t, secondChannelID, "task_public_7", "upstream_b_1")
secondChannelSecond := seedPollingTask(t, secondChannelID, "task_public_8", "upstream_b_2")
adaptor := &taskPollingFetchAdaptor{}
previousFactory := GetTaskAdaptorFunc
GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor { return adaptor }
t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory })
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
defer cancel()
err := UpdateVideoTasks(ctx, constant.TaskPlatform("kling"), map[int][]string{
firstChannelID: {
firstChannelFirst.GetUpstreamTaskID(),
firstChannelSecond.GetUpstreamTaskID(),
},
secondChannelID: {
secondChannelFirst.GetUpstreamTaskID(),
secondChannelSecond.GetUpstreamTaskID(),
},
}, map[string]*model.Task{
firstChannelFirst.GetUpstreamTaskID(): firstChannelFirst,
firstChannelSecond.GetUpstreamTaskID(): firstChannelSecond,
secondChannelFirst.GetUpstreamTaskID(): secondChannelFirst,
secondChannelSecond.GetUpstreamTaskID(): secondChannelSecond,
})
require.ErrorIs(t, err, context.DeadlineExceeded)
assert.ElementsMatch(t, []string{"upstream_a_1", "upstream_b_1"}, adaptor.fetchedTaskIDs())
}
func TestUpdateVideoTasksSlowChannelDoesNotBlockOtherChannels(t *testing.T) {
truncate(t)
const slowChannelID = 251
const fastChannelID = 252
seedTaskPollingChannel(t, slowChannelID, false)
seedTaskPollingChannel(t, fastChannelID, true)
slowTask := seedPollingTask(t, slowChannelID, "task_public_slow", "upstream_slow_1")
fastFirst := seedPollingTask(t, fastChannelID, "task_public_fast_1", "upstream_fast_parallel_1")
fastSecond := seedPollingTask(t, fastChannelID, "task_public_fast_2", "upstream_fast_parallel_2")
adaptor := &taskPollingFetchAdaptor{
fetched: make(chan string, 4),
blockTaskID: slowTask.GetUpstreamTaskID(),
blockStarted: make(chan struct{}),
releaseBlock: make(chan struct{}),
}
var releaseOnce sync.Once
releaseBlockedTask := func() {
releaseOnce.Do(func() {
close(adaptor.releaseBlock)
})
}
t.Cleanup(releaseBlockedTask)
previousFactory := GetTaskAdaptorFunc
GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor { return adaptor }
t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory })
errCh := make(chan error, 1)
gopool.Go(func() {
errCh <- UpdateVideoTasks(context.Background(), constant.TaskPlatform("kling"), map[int][]string{
slowChannelID: {
slowTask.GetUpstreamTaskID(),
},
fastChannelID: {
fastFirst.GetUpstreamTaskID(),
fastSecond.GetUpstreamTaskID(),
},
}, map[string]*model.Task{
slowTask.GetUpstreamTaskID(): slowTask,
fastFirst.GetUpstreamTaskID(): fastFirst,
fastSecond.GetUpstreamTaskID(): fastSecond,
})
})
select {
case <-adaptor.blockStarted:
case <-time.After(500 * time.Millisecond):
t.Fatal("slow channel did not start blocking")
}
require.Eventually(t, func() bool {
fetchedTaskIDs := adaptor.fetchedTaskIDs()
return len(fetchedTaskIDs) == 2 &&
fetchedTaskIDs[0] == fastFirst.GetUpstreamTaskID() &&
fetchedTaskIDs[1] == fastSecond.GetUpstreamTaskID()
}, 500*time.Millisecond, 10*time.Millisecond)
releaseBlockedTask()
require.NoError(t, <-errCh)
assert.ElementsMatch(t, []string{
slowTask.GetUpstreamTaskID(),
fastFirst.GetUpstreamTaskID(),
fastSecond.GetUpstreamTaskID(),
}, adaptor.fetchedTaskIDs())
}
func TestUpdateVideoTasksMixedChannelSleepSettings(t *testing.T) {
truncate(t)
const sleepyChannelID = 301
const fastChannelID = 302
seedTaskPollingChannel(t, sleepyChannelID, false)
seedTaskPollingChannel(t, fastChannelID, true)
sleepyFirst := seedPollingTask(t, sleepyChannelID, "task_public_9", "upstream_sleepy_1")
sleepySecond := seedPollingTask(t, sleepyChannelID, "task_public_10", "upstream_sleepy_2")
fastFirst := seedPollingTask(t, fastChannelID, "task_public_11", "upstream_fast_1")
fastSecond := seedPollingTask(t, fastChannelID, "task_public_12", "upstream_fast_2")
adaptor := &taskPollingFetchAdaptor{}
previousFactory := GetTaskAdaptorFunc
GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor { return adaptor }
t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory })
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()
err := UpdateVideoTasks(ctx, constant.TaskPlatform("kling"), map[int][]string{
sleepyChannelID: {
sleepyFirst.GetUpstreamTaskID(),
sleepySecond.GetUpstreamTaskID(),
},
fastChannelID: {
fastFirst.GetUpstreamTaskID(),
fastSecond.GetUpstreamTaskID(),
},
}, map[string]*model.Task{
sleepyFirst.GetUpstreamTaskID(): sleepyFirst,
sleepySecond.GetUpstreamTaskID(): sleepySecond,
fastFirst.GetUpstreamTaskID(): fastFirst,
fastSecond.GetUpstreamTaskID(): fastSecond,
})
require.ErrorIs(t, err, context.DeadlineExceeded)
assert.ElementsMatch(t, []string{"upstream_sleepy_1", "upstream_fast_1", "upstream_fast_2"}, adaptor.fetchedTaskIDs())
}
func TestUpdateSunoTasksStalePollsRefundExactlyOnce(t *testing.T) {
truncate(t)
const userID, tokenID, channelID = 401, 401, 401
const initialUserQuota, initialTokenQuota, taskQuota = 10_000, 6_000, 2_500
const publicTaskID, upstreamTaskID = "suno_public_refund_once", "suno_upstream_refund_once"
seedUser(t, userID, initialUserQuota)
seedToken(t, tokenID, userID, "sk-suno-refund-once", initialTokenQuota)
baseURL := "https://suno.invalid"
require.NoError(t, model.DB.Create(&model.Channel{
Id: channelID,
Type: constant.ChannelTypeSunoAPI,
Name: "suno_refund_once",
Key: "sk-suno-channel",
Status: common.ChannelStatusEnabled,
BaseURL: &baseURL,
}).Error)
task := makeTask(userID, channelID, taskQuota, tokenID, BillingSourceWallet, 0)
task.TaskID = publicTaskID
task.Platform = constant.TaskPlatformSuno
task.Status = model.TaskStatusInProgress
task.Progress = "50%"
task.SubmitTime = time.Now().Unix()
task.PrivateData.UpstreamTaskID = upstreamTaskID
require.NoError(t, model.DB.Create(task).Error)
var firstPollTask model.Task
var staleSecondPollTask model.Task
require.NoError(t, model.DB.First(&firstPollTask, task.ID).Error)
require.NoError(t, model.DB.First(&staleSecondPollTask, task.ID).Error)
adaptor := &sunoFailurePollingAdaptor{failReason: "upstream failed"}
previousFactory := GetTaskAdaptorFunc
GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor { return adaptor }
t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory })
require.NoError(t, updateSunoTasks(context.Background(), channelID, []string{upstreamTaskID}, map[string]*model.Task{
upstreamTaskID: &firstPollTask,
}))
require.NoError(t, updateSunoTasks(context.Background(), channelID, []string{upstreamTaskID}, map[string]*model.Task{
upstreamTaskID: &staleSecondPollTask,
}))
var reloaded model.Task
require.NoError(t, model.DB.First(&reloaded, task.ID).Error)
assert.EqualValues(t, model.TaskStatusFailure, reloaded.Status)
assert.Zero(t, reloaded.Quota)
assert.Equal(t, initialUserQuota+taskQuota, getUserQuota(t, userID))
assert.Equal(t, initialTokenQuota+taskQuota, getTokenRemainQuota(t, tokenID))
assert.Equal(t, int64(1), countLogs(t))
}
func TestRunTaskPollingOnceDoesNotRefundHistoricalFailedTask(t *testing.T) {
truncate(t)
const userID, initialQuota, taskQuota = 402, 10_000, 1_200
seedUser(t, userID, initialQuota)
task := makeTask(userID, 0, taskQuota, 0, BillingSourceWallet, 0)
task.TaskID = "historical_failed_already_refunded"
task.Status = model.TaskStatusFailure
task.Progress = "100%"
task.SubmitTime = time.Now().Add(-90 * 24 * time.Hour).Unix()
task.UpdatedAt = time.Now().Add(-time.Minute).Unix()
require.NoError(t, model.DB.Create(task).Error)
previousFactory := GetTaskAdaptorFunc
GetTaskAdaptorFunc = func(constant.TaskPlatform) TaskPollingAdaptor {
return &taskPollingFetchAdaptor{}
}
t.Cleanup(func() { GetTaskAdaptorFunc = previousFactory })
summary := RunTaskPollingOnce(context.Background(), nil)
assert.Zero(t, summary.UnfinishedTasks)
assert.Equal(t, initialQuota, getUserQuota(t, userID))
assert.Equal(t, taskQuota, getTaskQuota(t, task.ID))
assert.Equal(t, int64(0), countLogs(t))
}
func TestSweepTimedOutTasksHonorsRefundRolloutBoundary(t *testing.T) {
truncate(t)
const (
userID = 403
initialQuota = 10_000
legacyTaskQuota = 1_800
modernTaskQuota = 1_200
)
seedUser(t, userID, initialQuota)
legacyTask := makeTask(userID, 0, legacyTaskQuota, 0, BillingSourceWallet, 0)
legacyTask.TaskID = "legacy_timeout_without_refund"
legacyTask.Progress = "50%"
legacyTask.SubmitTime = 1771718399 // 2026-02-21 23:59:59 UTC
require.NoError(t, model.DB.Create(legacyTask).Error)
modernTask := makeTask(userID, 0, modernTaskQuota, 0, BillingSourceWallet, 0)
modernTask.TaskID = "modern_timeout_with_refund"
modernTask.Progress = "50%"
modernTask.SubmitTime = 1771718400 // 2026-02-22 00:00:00 UTC
require.NoError(t, model.DB.Create(modernTask).Error)
previousTimeout := constant.TaskTimeoutMinutes
constant.TaskTimeoutMinutes = 1
t.Cleanup(func() { constant.TaskTimeoutMinutes = previousTimeout })
sweepTimedOutTasks(context.Background())
var reloadedLegacy model.Task
var reloadedModern model.Task
require.NoError(t, model.DB.First(&reloadedLegacy, legacyTask.ID).Error)
require.NoError(t, model.DB.First(&reloadedModern, modernTask.ID).Error)
assert.EqualValues(t, model.TaskStatusFailure, reloadedLegacy.Status)
assert.EqualValues(t, model.TaskStatusFailure, reloadedModern.Status)
assert.Zero(t, reloadedLegacy.Quota)
assert.Zero(t, reloadedModern.Quota)
assert.Contains(t, reloadedLegacy.FailReason, "旧系统遗留任务")
assert.Contains(t, reloadedModern.FailReason, "任务超时")
assert.Equal(t, initialQuota+modernTaskQuota, getUserQuota(t, userID))
assert.Equal(t, int64(1), countLogs(t))
}