Commit 9df450fe by CaIon

feat(task): give polling hooks a real query context, host HTTP classification,…

feat(task): give polling hooks a real query context, host HTTP classification, and bounded poll failures

Plugin polling hooks previously ran against a hollow context: parseTaskResult
and parseBatchResult received {} / nil, buildQueryRequest received a
{task_id, action} map under the misleading name requestBody, and batch hooks
saw only bare task ids. The per-task poller also never looked at the upstream
HTTP status, and every built-in plugin papered over unrecognized bodies with
`|| "IN_PROGRESS"`, so a 404, a revoked key, or a shape the plugin did not
know would sit in IN_PROGRESS for the full 24h TASK_TIMEOUT_MINUTES while
holding the user's pre-charged quota.

Contract (docs/plugin-api v1.d.ts, v1.md, v1.schema.json):
- TaskQueryContext is declared separately from DriverContext and rebuilt from
  the persisted Task row: taskId, publicTaskId, action, model, upstreamModel,
  baseUrl, apiKey, authHeader, auth, data, state. Query-side requestBody is
  removed; the original request is not persisted and hooks that need a
  request-derived value must save it into state at submit time.
- parseTaskResult / parseBatchResult receive a third {status, headers}
  argument. Batch hooks receive tasks[] with one TaskQueryContext per task.
- NormalizedTaskResult accepts status "UNKNOWN" meaning "I do not recognize
  this body". Falling back to IN_PROGRESS for unknown shapes is forbidden;
  `plugin lint` warns on the literal.
- parseSubmitResponse / parseTaskResult / parseBatchResult may return `state`.
  Task.Data remains a per-round snapshot overwritten on every valid parse;
  state is plugin-owned, persisted in TaskPrivateData.PluginState, preserved
  when a hook omits it, byte-capped like taskData, and never exposed through
  presenter views.

Host (service/task_polling.go, relay/channel/task/jsplugin/adaptor.go):
- TaskPollingAdaptor / BatchTaskPollingAdaptor take *model.Task and the
  *http.Response so the adaptor can build the full context; jsplugin is the
  only implementation.
- HTTP classification before the plugin sees the body: 2xx -> plugin;
  404/410 -> FAILURE and refund; 401/403 -> poll failure plus a channel-scoped
  warning, no auto-disable; 429/5xx/transport -> poll failure; other 4xx ->
  plugin with the status visible, counted as unrecognized if the plugin still
  reports a non-terminal state.
- TaskPrivateData.PollFailures counts consecutive poll failures (transient
  HTTP, auth, transport, hook error, UNKNOWN). It is persisted through the
  existing UpdateWithStatus CAS so a concurrent terminal transition on another
  instance is never clobbered, and reset on any valid 2xx non-terminal parse.
  Reaching TASK_POLL_MAX_FAILURES (default 20, <= 0 disables) fails the task
  with the last classification and HTTP code in fail_reason and runs the
  existing settle/refund chain exactly once. sweepTimedOutTasks and its
  1440-minute default are unchanged as the outer backstop.
- Unrecognized bodies are logged at WARN with a bounded redacted copy since
  Task.Data is intentionally not overwritten on that path.

Plugins (all ten bumped one patch version):
- jimeng persists the outbound req_key in state and reads it back in
  buildQueryRequest, replacing dead reads of ctx.data / ctx.requestBody that
  never resolved.
- sunoapi batch hooks read tasks[] instead of the removed requestBody.
- hailuo treats base_resp.status_code != 0 as FAILURE before the status table.
- kling, vidu, sora, alibaba, doubao, hailuo, jimeng return UNKNOWN with the
  raw upstream status in reason on table miss.
- google and vertex-ai treat a missing `done` as in-progress: Google
  long-running operations omit proto3 default fields, so a running Veo
  operation has no `done` key at all. Only a body without an operation name is
  UNKNOWN. plugins/veo_poll_test.go locks this so the poll-failure cutoff can
  never fail a rendering Veo task.

Tests cover the classification table end to end against a real DB (404
immediate refund, 429xN refund, 401 increments without status change, 2xx
reset, UNKNOWN increments, state preserved vs replaced, PollFailures survives
the CAS write), the query-context shape, UNKNOWN on unrecognized bodies, and
the absence of PluginState/PollFailures from TaskView. Controller tests derive
the kling factory version from the embedded manifest instead of hardcoding it.
parent 9f506dd7
......@@ -56,6 +56,10 @@
# 任务和功能配置
# 更新任务启用
# UPDATE_TASK=true
# 异步任务硬超时(分钟),按提交时间计算,超时未完成的任务标记失败并退款;0 表示禁用
# TASK_TIMEOUT_MINUTES=1440
# 异步任务连续轮询失败阈值(上游 429/5xx/401/403、网络错误、无法识别的响应),达到后任务标记失败并退款;正常轮询成功一次即归零
# TASK_POLL_MAX_FAILURES=20
# 对话超时设置
# 所有请求超时时间,单位秒,默认为0,表示不限制
......
......@@ -202,6 +202,8 @@ func initConstantEnv() {
constant.TaskQueryLimit = GetEnvOrDefault("TASK_QUERY_LIMIT", 1000)
// 异步任务超时时间(分钟),超过此时间未完成的任务将被标记为失败并退款。0 表示禁用。
constant.TaskTimeoutMinutes = GetEnvOrDefault("TASK_TIMEOUT_MINUTES", 1440)
// Consecutive unrecognized/transient poll failures before the task is failed and refunded.
constant.TaskPollMaxFailures = GetEnvOrDefault("TASK_POLL_MAX_FAILURES", 20)
// 声明式任务协议桥只观察数据库;这些值控制一次客户端观察连接,
// 不改变后台轮询或结算生命周期。
constant.TaskPluginProtocolTimeoutSeconds = GetEnvOrDefault("TASK_PLUGIN_PROTOCOL_TIMEOUT_SECONDS", 600)
......
......@@ -18,6 +18,7 @@ var GenerateDefaultToken bool
var ErrorLogEnabled bool
var TaskQueryLimit int
var TaskTimeoutMinutes int
var TaskPollMaxFailures = 20
var TaskPluginProtocolTimeoutSeconds int
var TaskPluginProtocolTickMilliseconds int
var TaskPluginProtocolTickJitterMilliseconds int
......
......@@ -366,14 +366,14 @@ type terminalSettlementPollingAdaptor struct {
func (a *terminalSettlementPollingAdaptor) Init(*relaycommon.RelayInfo) {}
func (a *terminalSettlementPollingAdaptor) FetchTask(string, string, map[string]any, string) (*http.Response, error) {
func (a *terminalSettlementPollingAdaptor) FetchTask(string, string, *model.Task, string) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{}`)),
}, nil
}
func (a *terminalSettlementPollingAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) {
func (a *terminalSettlementPollingAdaptor) ParseTaskResult(*model.Task, *http.Response, []byte) (*relaycommon.TaskInfo, error) {
return &relaycommon.TaskInfo{
Status: model.TaskStatusSuccess,
Progress: "100%",
......
......@@ -744,6 +744,9 @@ func executeTaskSubmissionWith(
}
task.Quota = result.Quota
task.Data = result.TaskData
if len(result.PluginState) > 0 {
task.PrivateData.PluginState = result.PluginState
}
task.Action = relayInfo.Action
if immediate := result.Immediate; immediate != nil {
task.Status = model.TaskStatus(immediate.Status)
......
......@@ -5,6 +5,7 @@ import (
"fmt"
"net/http"
"net/http/httptest"
"regexp"
"strings"
"testing"
......@@ -117,6 +118,17 @@ func TestDisableThirdPartyPluginSupportsCascadeAndForce(t *testing.T) {
assert.Equal(t, common.ChannelStatusManuallyDisabled, updated.Status)
}
// klingFactoryVersion returns the version declared in the embedded kling factory
// manifest so tests do not hardcode a value that moves with every plugin release.
func klingFactoryVersion(t *testing.T) string {
t.Helper()
factorySource, err := plugins.Source("kling")
require.NoError(t, err)
match := regexp.MustCompile(`version:\s*"([^"]+)"`).FindStringSubmatch(factorySource)
require.Len(t, match, 2, "kling factory manifest must declare a version")
return match[1]
}
func setupTaskPluginFactoryDisableTest(t *testing.T) {
t.Helper()
setupTaskPluginControllerTest(t)
......@@ -240,7 +252,9 @@ func TestDisableFactoryOverrideRowKeepsEnabledFlagPath(t *testing.T) {
setupTaskPluginFactoryDisableTest(t)
factorySource, err := plugins.Source("kling")
require.NoError(t, err)
overrideSource := strings.Replace(factorySource, `version: "1.0.0"`, `version: "1.0.0-test-factory-status"`, 1)
factoryVersion := klingFactoryVersion(t)
overrideSource := strings.Replace(factorySource, `version: "`+factoryVersion+`"`, `version: "`+factoryVersion+`-test-factory-status"`, 1)
require.NotEqual(t, factorySource, overrideSource, "factory version marker must be found in kling source")
loaded, err := jsplugin.DefaultRegistry.Register(overrideSource, jsplugin.Options{})
require.NoError(t, err)
t.Cleanup(func() { jsplugin.DefaultRegistry.Unregister("kling") })
......@@ -266,7 +280,7 @@ func TestDisableFactoryOverrideRowKeepsEnabledFlagPath(t *testing.T) {
assert.True(t, taskPluginOptionsHasKey(t, "kling"))
got, ok := jsplugin.DefaultRegistry.Get("kling")
require.True(t, ok)
assert.Equal(t, "1.0.0", got.Meta.Version)
assert.Equal(t, factoryVersion, got.Meta.Version)
}
func TestListTaskPluginsIncludesFactoryWithoutDatabaseRows(t *testing.T) {
......@@ -358,7 +372,9 @@ func TestListTaskPluginsShowsDisabledFallbackWhenOverridesAreDisabled(t *testing
setupTaskPluginControllerTest(t)
factorySource, err := plugins.Source("kling")
require.NoError(t, err)
overrideSource := strings.Replace(factorySource, `version: "1.0.0"`, `version: "1.0.0-test-disabled-override"`, 1)
factoryVersion := klingFactoryVersion(t)
overrideSource := strings.Replace(factorySource, `version: "`+factoryVersion+`"`, `version: "`+factoryVersion+`-test-disabled-override"`, 1)
require.NotEqual(t, factorySource, overrideSource, "factory version marker must be found in kling source")
loaded, err := jsplugin.DefaultRegistry.Register(overrideSource, jsplugin.Options{})
require.NoError(t, err)
plugin := model.TaskPlugin{
......@@ -399,7 +415,9 @@ func TestDeleteActiveOverrideFallsBackToFactoryAndDeletesRecord(t *testing.T) {
setupTaskPluginControllerTest(t)
factorySource, err := plugins.Source("kling")
require.NoError(t, err)
overrideSource := strings.Replace(factorySource, `version: "1.0.0"`, `version: "1.0.0-test-override"`, 1)
factoryVersion := klingFactoryVersion(t)
overrideSource := strings.Replace(factorySource, `version: "`+factoryVersion+`"`, `version: "`+factoryVersion+`-test-override"`, 1)
require.NotEqual(t, factorySource, overrideSource, "factory version marker must be found in kling source")
loaded, err := jsplugin.DefaultRegistry.Register(overrideSource, jsplugin.Options{Key: "kling", Version: "test-override"})
require.NoError(t, err)
t.Cleanup(func() { jsplugin.DefaultRegistry.Unregister("kling") })
......
......@@ -25,6 +25,9 @@ export type UsageExample = {label: string; facts: Readonly<Record<string, string
export interface Meta {apiVersion: 1; key: string; name: string; icon?: string; description?: LocalizedText; version: string; author: {name: string; url?: string}; channelTypes?: readonly number[]; models: readonly string[]; fetchMode: "per_task" | "batch"; allowedHosts?: readonly string[]; routes?: readonly NativeRoute[]; protocols?: readonly ProtocolClaim[]; usageSchema?: Readonly<Record<string, UsageFieldSchema>>; usageExamples?: readonly UsageExample[]; auth?: "none" | "api_key" | "vertex_oauth" | {type: "none" | "api_key" | "oauth2_jwt"}}
export interface TaskView {task_id: string; status: string; progress?: string; fail_reason?: string; created_at?: number; updated_at?: number; data?: unknown; properties?: Record<string, unknown>}
export interface DriverContext {requestBody: unknown; requestHeaders: Readonly<Record<string, string>>; action: string; model: string; upstreamModel: string; baseUrl: string; apiKey?: string; authHeader: string; files: readonly FileReference[]; publicTaskId: string; originTasks?: readonly {taskId: string; upstreamTaskId: string; action: string; status: string; data: unknown}[]}
export interface TaskQueryContext {taskId: string; publicTaskId: string; action: string; model: string; upstreamModel: string; baseUrl: string; apiKey?: string; authHeader: string; auth?: unknown; data: unknown; state: unknown}
export interface BatchQueryContext {baseUrl: string; apiKey?: string; authHeader: string; auth?: unknown; tasks: readonly TaskQueryContext[]}
export type HookHTTPResponse = {readonly status: number; readonly headers: Readonly<Record<string, string>>}
export interface RequestDescriptor {url: string; method?: string; headers?: Record<string, string>; /** JSON body may contain FilePlaceholder objects at any depth; the host replaces each with a Base64 or data-URL string. */ body?: unknown; credentialless?: boolean; action?: string; model?: string; rewriteModel?: string; bodyType?: "json" | "multipart"; parts?: readonly {name: string; value?: unknown; fileRef?: string; filename?: string}[]}
export interface UpstreamResponse {statusCode: number; headers: Readonly<Record<string, readonly string[]>>; body: unknown}
export interface NormalizedTaskResult {taskId?: string; status: "NOT_START" | "SUBMITTED" | "QUEUED" | "IN_PROGRESS" | "SUCCESS" | "FAILURE" | "UNKNOWN"; progress?: string; reason?: string; url?: string; remoteUrl?: string; completionTokens?: number; totalTokens?: number}
......@@ -36,13 +39,13 @@ export declare const protocols: {
openai_video?: {decodeRequest(ctx: ProtocolDecodeContext): SubmitIntent; render(ctx: unknown, task: TaskView): unknown};
};
export declare function buildSubmitRequest(ctx: DriverContext): RequestDescriptor;
export declare function parseSubmitResponse(ctx: DriverContext, response: UpstreamResponse): {taskId: string; taskData?: unknown; immediate?: NormalizedTaskResult};
export declare function buildQueryRequest(ctx: DriverContext & {taskId: string}): RequestDescriptor;
export declare function buildBatchQueryRequest(ctx: DriverContext, taskIds: readonly string[]): RequestDescriptor;
export declare function parseTaskResult(ctx: DriverContext, body: unknown): NormalizedTaskResult;
export declare function parseBatchResult(ctx: DriverContext, body: unknown): readonly (NormalizedTaskResult & {taskId: string; data?: unknown})[];
export declare function parseSubmitResponse(ctx: DriverContext, response: UpstreamResponse): {taskId: string; taskData?: unknown; immediate?: NormalizedTaskResult; state?: unknown};
export declare function buildQueryRequest(ctx: TaskQueryContext): RequestDescriptor;
export declare function buildBatchQueryRequest(ctx: BatchQueryContext, tasks: readonly TaskQueryContext[]): RequestDescriptor;
export declare function parseTaskResult(ctx: TaskQueryContext, body: unknown, response: HookHTTPResponse): NormalizedTaskResult;
export declare function parseBatchResult(ctx: BatchQueryContext, body: unknown, response: HookHTTPResponse): readonly (NormalizedTaskResult & {taskId: string; data?: unknown; state?: unknown})[];
export declare function extractUsage(ctx: DriverContext & {usagePurpose?: "facts" | "billing_ratios"}): Readonly<Record<string, string | number | boolean>> | null;
export declare function extractUsageOnSubmit(ctx: DriverContext, taskData: unknown): Readonly<Record<string, string | number | boolean>> | null;
export declare function extractUsageOnComplete(task: TaskView, result: NormalizedTaskResult, data: unknown): Readonly<Record<string, string | number | boolean>> | null;
export declare function listArtifacts(task: {taskId: string; status: string; action: string; data: unknown; producerVersion: string}): readonly TaskArtifact[];
export declare function buildContentRequest(ctx: DriverContext & {artifactKey: string; data: unknown; upstreamTaskId: string; clientRequest: {method: "GET" | "HEAD"; headers: Readonly<Record<string, string>>}}): RequestDescriptor;
export declare function buildContentRequest(ctx: DriverContext & {artifactKey: string; data: unknown; state?: unknown; upstreamTaskId: string; clientRequest: {method: "GET" | "HEAD"; headers: Readonly<Record<string, string>>}}): RequestDescriptor;
......@@ -32,7 +32,7 @@ Each `protocols` entry claims a host protocol. A protocol that defines modes mus
Enabled uploads pre-flight the candidate against the live routing generation and reject the first channel-type, native-route, or protocol-model conflict (the error names the counterpart plugin). Set `force: true` or `enabled: false` to store the plugin anyway.
`endpoints`, `routes[].renderer`, global `resolveRequest`, global `renderError`, and global `renderers` are rejected. `parseSubmitResponse` returns only `{taskId, taskData}` (plus the documented lifecycle fields); `clientResponse` is rejected.
`endpoints`, `routes[].renderer`, global `resolveRequest`, global `renderError`, and global `renderers` are rejected. `parseSubmitResponse` returns only `{taskId, taskData, immediate?, state?}`; `clientResponse` is rejected.
`icon` is an optional LobeHub icon name string (for example `Sora.Color`). The values `text` and `text:<label>` request a generated text avatar instead (label defaults to the first two characters of `name`). It is display-only and does not participate in routing, billing, or admission beyond type and length checks.
......@@ -140,3 +140,38 @@ Protocol media uses host-injected `ctx.artifacts[key].url`. Provider URLs from `
The persisted field remains `task.data`; there is no `task.raw` alias. Driver hooks (`buildSubmitRequest`, `parseSubmitResponse`, query/result, usage, artifact, and content hooks) stay flat and must not branch on the client path or protocol.
`ctx.model` is the billing and display identity (the origin name the client sent, including a channel-mapping alias). `ctx.upstreamModel` is the machine identity after channel `model_mapping`. Rate tables and model-keyed usage facts must use `ctx.upstreamModel || ctx.model`. Decode and render hooks that echo the client model must keep `ctx.model`. `buildSubmitRequest` must not set descriptor top-level `model` on a mapped pin; the host requires the plugin to echo the alias verbatim. Background polling has no relay info, so query hooks receive both identities from the persisted task properties, and `ctx.upstreamModel` falls back to `ctx.model` when the task was submitted without a channel mapping.
## Polling contract
Query and parse hooks use `TaskQueryContext`, not `DriverContext`. The host rebuilds that context from the persisted task row. There is no query-side `requestBody`.
| Field | Source |
|-------|--------|
| `taskId` | Upstream task id (`PrivateData.UpstreamTaskID`, else `TaskID`) |
| `publicTaskId` | Gateway task id |
| `action` | Normalized persisted action |
| `model` | `Properties.OriginModelName` |
| `upstreamModel` | `Properties.UpstreamModelName`, falling back to `model` |
| `baseUrl` / `apiKey` / `authHeader` / `auth` | Channel credentials |
| `data` | Current `Task.Data` snapshot |
| `state` | Plugin-owned `PrivateData.PluginState` |
`Task.Data` is the latest upstream response snapshot for presenters and artifacts. The host overwrites it on every successful parse. Values that must survive across poll rounds belong in `state`.
`parseSubmitResponse`, `parseTaskResult`, and each `parseBatchResult` item may return optional `state`. The host writes it only when the hook returns it. Omitting `state` preserves the previous value. Oversized state is rejected with a warning, not truncated.
`buildBatchQueryRequest(ctx, tasks)` and `parseBatchResult` receive `tasks: TaskQueryContext[]`. `parseTaskResult` / `parseBatchResult` also receive `{status, headers}` for the upstream HTTP response.
`status: "UNKNOWN"` means the plugin does not recognize the response. Do not write `|| "IN_PROGRESS"` (or equivalent) for a missing table entry. The host treats `UNKNOWN`, hook errors, empty status, and unrecognized status strings as consecutive poll failures.
The host classifies the HTTP status before trusting a non-terminal parse:
| Upstream HTTP | Host action |
|---------------|-------------|
| 2xx | Call the parse hook |
| 404 / 410 | Immediate `FAILURE` and refund |
| 401 / 403 | Leave task status unchanged; increment `PollFailures`; `LogWarn` with channel id. Channels are not auto-disabled. |
| 429 / 5xx / transport error | Increment `PollFailures` |
| Other 4xx | Call the parse hook with `response.status`. A still-non-terminal result is unrecognized and increments `PollFailures`. |
A valid 2xx non-terminal parse resets `PollFailures` to 0. After `TASK_POLL_MAX_FAILURES` (default 20) consecutive failures the task becomes `FAILURE` and follows the existing refund chain. The 24h `TASK_TIMEOUT_MINUTES` sweep remains the outer deadline.
......@@ -41,6 +41,45 @@
"maxBytes": {"type": "integer", "exclusiveMinimum": 0}
}
},
"route": {"type": "object", "additionalProperties": false, "required": ["method", "path", "type", "render"], "properties": {"method": {"enum": ["GET", "POST", "PUT", "PATCH", "DELETE"]}, "path": {"type": "string", "pattern": "^/"}, "type": {"enum": ["submit", "query", "dynamic"]}, "action": {"type": "string"}, "taskIdParam": {"type": "string"}, "decode": {"type": "string"}, "render": {"type": "string"}, "models": {"type": "array", "minItems": 1, "uniqueItems": true, "items": {"type": "string", "minLength": 1}}}, "allOf": [{"if": {"properties": {"type": {"const": "query"}}}, "then": {"allOf": [{"not": {"required": ["decode"]}}, {"not": {"required": ["models"]}}]}}, {"if": {"properties": {"type": {"enum": ["submit", "dynamic"]}}}, "then": {"required": ["decode"]}}]}
"route": {"type": "object", "additionalProperties": false, "required": ["method", "path", "type", "render"], "properties": {"method": {"enum": ["GET", "POST", "PUT", "PATCH", "DELETE"]}, "path": {"type": "string", "pattern": "^/"}, "type": {"enum": ["submit", "query", "dynamic"]}, "action": {"type": "string"}, "taskIdParam": {"type": "string"}, "decode": {"type": "string"}, "render": {"type": "string"}, "models": {"type": "array", "minItems": 1, "uniqueItems": true, "items": {"type": "string", "minLength": 1}}}, "allOf": [{"if": {"properties": {"type": {"const": "query"}}}, "then": {"allOf": [{"not": {"required": ["decode"]}}, {"not": {"required": ["models"]}}]}}, {"if": {"properties": {"type": {"enum": ["submit", "dynamic"]}}}, "then": {"required": ["decode"]}}]},
"taskQueryContext": {
"type": "object",
"additionalProperties": false,
"required": ["taskId", "publicTaskId", "action", "model", "upstreamModel", "baseUrl", "authHeader", "data", "state"],
"properties": {
"taskId": {"type": "string"},
"publicTaskId": {"type": "string"},
"action": {"type": "string"},
"model": {"type": "string"},
"upstreamModel": {"type": "string"},
"baseUrl": {"type": "string"},
"apiKey": {"type": "string"},
"authHeader": {"type": "string"},
"auth": true,
"data": true,
"state": true
}
},
"batchQueryContext": {
"type": "object",
"additionalProperties": false,
"required": ["baseUrl", "authHeader", "tasks"],
"properties": {
"baseUrl": {"type": "string"},
"apiKey": {"type": "string"},
"authHeader": {"type": "string"},
"auth": true,
"tasks": {"type": "array", "items": {"$ref": "#/$defs/taskQueryContext"}}
}
},
"hookHTTPResponse": {
"type": "object",
"additionalProperties": false,
"required": ["status", "headers"],
"properties": {
"status": {"type": "integer"},
"headers": {"type": "object", "additionalProperties": {"type": "string"}}
}
}
}
}
......@@ -140,9 +140,7 @@ func (o *LogOther) toMap() map[string]any {
return result
}
for key, value := range o.public {
result[key] = value
}
maps.Copy(result, o.public)
if adminInfo := copyLogOtherMap(o.adminInfo); len(adminInfo) > 0 {
result[logOtherAdminInfoKey] = adminInfo
}
......
......@@ -126,6 +126,11 @@ type TaskPrivateData struct {
// disconnect regardless; this only echoes the protocol-level request
// attribute back on retrieval snapshots.
ResponsesBackground bool `json:"responses_background,omitempty"`
// PluginState is plugin-owned cross-round data. Unlike Task.Data it is
// only replaced when a hook explicitly returns state.
PluginState json.RawMessage `json:"plugin_state,omitempty"`
// PollFailures counts consecutive unrecognized or transient poll outcomes.
PollFailures int `json:"poll_failures,omitempty"`
}
type TaskExecutionSnapshot struct {
......@@ -194,7 +199,10 @@ func (p *TaskPrivateData) Scan(val interface{}) error {
}
func (p TaskPrivateData) Value() (driver.Value, error) {
if (p == TaskPrivateData{}) {
if p.Key == "" && p.UpstreamTaskID == "" && p.ResultURL == "" &&
p.Execution == nil && p.BillingSource == "" && p.SubscriptionId == 0 &&
p.TokenId == 0 && p.NodeName == "" && p.BillingContext == nil &&
!p.ResponsesBackground && len(p.PluginState) == 0 && p.PollFailures == 0 {
return nil, nil
}
// 同 Properties.Value:string 避免 PG simple protocol 的 bytea 编码。
......@@ -466,13 +474,15 @@ func (Task *Task) InsertWithContext(ctx context.Context) error {
}
type taskSnapshot struct {
Status TaskStatus
Progress string
StartTime int64
FinishTime int64
FailReason string
ResultURL string
Data json.RawMessage
Status TaskStatus
Progress string
StartTime int64
FinishTime int64
FailReason string
ResultURL string
Data json.RawMessage
PluginState json.RawMessage
PollFailures int
}
func (s taskSnapshot) Equal(other taskSnapshot) bool {
......@@ -482,18 +492,22 @@ func (s taskSnapshot) Equal(other taskSnapshot) bool {
s.FinishTime == other.FinishTime &&
s.FailReason == other.FailReason &&
s.ResultURL == other.ResultURL &&
bytes.Equal(s.Data, other.Data)
bytes.Equal(s.Data, other.Data) &&
bytes.Equal(s.PluginState, other.PluginState) &&
s.PollFailures == other.PollFailures
}
func (t *Task) Snapshot() taskSnapshot {
return taskSnapshot{
Status: t.Status,
Progress: t.Progress,
StartTime: t.StartTime,
FinishTime: t.FinishTime,
FailReason: t.FailReason,
ResultURL: t.PrivateData.ResultURL,
Data: t.Data,
Status: t.Status,
Progress: t.Progress,
StartTime: t.StartTime,
FinishTime: t.FinishTime,
FailReason: t.FailReason,
ResultURL: t.PrivateData.ResultURL,
Data: t.Data,
PluginState: t.PrivateData.PluginState,
PollFailures: t.PrivateData.PollFailures,
}
}
......
......@@ -177,6 +177,29 @@ func TestSnapshotEqual_NilVsEmpty(t *testing.T) {
assert.True(t, a.Equal(b))
}
func TestSnapshotEqual_PluginStateAndPollFailures(t *testing.T) {
base := taskSnapshot{
Status: TaskStatusInProgress,
PluginState: json.RawMessage(`{"req_key":"a"}`),
PollFailures: 2,
}
assert.True(t, base.Equal(taskSnapshot{
Status: TaskStatusInProgress,
PluginState: json.RawMessage(`{"req_key":"a"}`),
PollFailures: 2,
}))
assert.False(t, base.Equal(taskSnapshot{
Status: TaskStatusInProgress,
PluginState: json.RawMessage(`{"req_key":"b"}`),
PollFailures: 2,
}))
assert.False(t, base.Equal(taskSnapshot{
Status: TaskStatusInProgress,
PluginState: json.RawMessage(`{"req_key":"a"}`),
PollFailures: 3,
}))
}
func TestSnapshot_Roundtrip(t *testing.T) {
task := &Task{
Status: TaskStatusInProgress,
......@@ -185,7 +208,9 @@ func TestSnapshot_Roundtrip(t *testing.T) {
FinishTime: 5678,
FailReason: "timeout",
PrivateData: TaskPrivateData{
ResultURL: "https://example.com/result.mp4",
ResultURL: "https://example.com/result.mp4",
PluginState: json.RawMessage(`{"req_key":"keep"}`),
PollFailures: 3,
},
Data: json.RawMessage(`{"model":"test-model"}`),
}
......@@ -197,6 +222,8 @@ func TestSnapshot_Roundtrip(t *testing.T) {
assert.Equal(t, task.FailReason, snap.FailReason)
assert.Equal(t, task.PrivateData.ResultURL, snap.ResultURL)
assert.JSONEq(t, string(task.Data), string(snap.Data))
assert.Equal(t, task.PrivateData.PluginState, snap.PluginState)
assert.Equal(t, task.PrivateData.PollFailures, snap.PollFailures)
}
// ---------------------------------------------------------------------------
......@@ -292,3 +319,30 @@ func TestUpdateWithStatus_ConcurrentWinner(t *testing.T) {
}
assert.Equal(t, 1, winCount, "exactly one goroutine should win the CAS")
}
func TestUpdateWithStatus_PersistsPluginStateAndPollFailures(t *testing.T) {
truncateTables(t)
task := &Task{
TaskID: "task_cas_plugin_state",
Status: TaskStatusInProgress,
Data: json.RawMessage(`{}`),
PrivateData: TaskPrivateData{
PluginState: json.RawMessage(`{"req_key":"old"}`),
PollFailures: 1,
},
}
insertTask(t, task)
task.PrivateData.PluginState = json.RawMessage(`{"req_key":"new"}`)
task.PrivateData.PollFailures = 4
won, err := task.UpdateWithStatus(TaskStatusInProgress)
require.NoError(t, err)
require.True(t, won)
var reloaded Task
require.NoError(t, DB.First(&reloaded, task.ID).Error)
assert.EqualValues(t, TaskStatusInProgress, reloaded.Status)
assert.JSONEq(t, `{"req_key":"new"}`, string(reloaded.PrivateData.PluginState))
assert.Equal(t, 4, reloaded.PrivateData.PollFailures)
}
......@@ -33,6 +33,7 @@ func RunCLI(args []string, stdout, stderr io.Writer) int {
fmt.Fprintf(stderr, "plugin lint failed: %v\n", compileErr)
return 1
}
warnParseTaskResultInProgressFallback(string(source), stderr)
fmt.Fprintf(stdout, "plugin %s@%s is valid\n", plugin.Meta.Key, plugin.Meta.Version)
return 0
case "test":
......@@ -57,3 +58,35 @@ func RunCLI(args []string, stdout, stderr io.Writer) int {
return 2
}
}
func warnParseTaskResultInProgressFallback(source string, stderr io.Writer) {
body := parseTaskResultFunctionBody(source)
if strings.Contains(body, `|| "IN_PROGRESS"`) || strings.Contains(body, `|| 'IN_PROGRESS'`) {
fmt.Fprintln(stderr, `warning: parseTaskResult uses || "IN_PROGRESS" fallback; return UNKNOWN for unrecognized statuses`)
}
}
func parseTaskResultFunctionBody(source string) string {
marker := strings.Index(source, "function parseTaskResult")
if marker < 0 {
return ""
}
brace := strings.Index(source[marker:], "{")
if brace < 0 {
return ""
}
start := marker + brace
depth := 0
for i := start; i < len(source); i++ {
switch source[i] {
case '{':
depth++
case '}':
depth--
if depth == 0 {
return source[start : i+1]
}
}
}
return ""
}
......@@ -28,6 +28,25 @@ func TestPluginCLI(t *testing.T) {
assert.Contains(t, stdout.String(), "1/1 cases")
}
func TestPluginCLIWarnsOnParseTaskResultInProgressFallback(t *testing.T) {
tempDir := t.TempDir()
pluginPath := filepath.Join(tempDir, "fallback.js")
require.NoError(t, os.WriteFile(pluginPath, []byte(`
export const meta = { apiVersion: 1, key: "fallback", name: "Fallback", version: "1.0.0", author: {name: "Test"}, models: ["m"], fetchMode: "per_task" };
export function buildSubmitRequest(ctx) { return {url: ctx.baseUrl}; }
export function parseSubmitResponse() { return {taskId: "task"}; }
export function buildQueryRequest(ctx) { return {url: ctx.baseUrl}; }
export function parseTaskResult(ctx, body) { return {status: statuses[body.status] || "IN_PROGRESS"}; }
const statuses = { done: "SUCCESS" };
`), 0o600))
var stdout bytes.Buffer
var stderr bytes.Buffer
assert.Equal(t, 0, RunCLI([]string{"lint", pluginPath}, &stdout, &stderr))
assert.Contains(t, stdout.String(), "plugin fallback@1.0.0 is valid")
assert.Contains(t, stderr.String(), `|| "IN_PROGRESS"`)
}
const cliFixturePluginSource = `
export const meta = { apiVersion: 1, key: "cli-fixture", name: "CLI Fixture", version: "1.0.0", author: {name: "Test"}, channelTypes: [1003], models: ["fixture-model"], fetchMode: "per_task" };
export function buildSubmitRequest(ctx) { return {url: ctx.baseUrl}; }
......
......@@ -346,6 +346,9 @@ func TestHailuoParseTaskResult(t *testing.T) {
{"H3 permanent query error", `{"type":"error","error":{"type":"authorized_error","message":"login failed","http_code":"401"}}`, "FAILURE", "", "login failed"},
{"legacy success", `{"task_id":"1","status":"Success","file_id":"f1","base_resp":{"status_code":0}}`, "SUCCESS", "", ""},
{"legacy processing", `{"task_id":"1","status":"Processing","base_resp":{"status_code":0}}`, "IN_PROGRESS", "", ""},
{"H3 unrecognized", `{"task":{"id":"1","status":"weird"}}`, "UNKNOWN", "", "unrecognized status: weird"},
{"legacy unrecognized", `{"task_id":"1","status":"Weird","base_resp":{"status_code":0}}`, "UNKNOWN", "", "unrecognized status: Weird"},
{"legacy base_resp failure", `{"task_id":"1","status":"Success","base_resp":{"status_code":1001,"status_msg":"upstream down"}}`, "FAILURE", "", "upstream down"},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
......@@ -451,7 +454,7 @@ func TestHailuoH3CompletionUsageFacts(t *testing.T) {
t.Run("polling adaptor carries actual facts into task settlement", func(t *testing.T) {
adaptor := taskplugin.New(plugin)
result, err := adaptor.ParseTaskResult([]byte(
result, err := adaptor.ParseTaskResult(&model.Task{}, &http.Response{StatusCode: http.StatusOK, Header: make(http.Header)}, []byte(
`{"task":{"id":"1","status":"succeeded","resolution":"2K","usage":{"output_seconds":5,"input_seconds":7.5,"input_image_count":6}}}`,
))
require.NoError(t, err)
......
package plugins_test
import "testing"
import (
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/pkg/jsplugin"
builtinplugins "github.com/QuantumNous/new-api/plugins"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestJimengResponsesProtocol(t *testing.T) {
testVideoResponsesProtocol(t, videoResponsesTestCase{
......@@ -25,3 +33,60 @@ func TestJimengResponsesProtocol(t *testing.T) {
wantVendorName: "jimeng",
})
}
func loadJimengPlugin(t *testing.T) *jsplugin.LoadedPlugin {
t.Helper()
source, err := builtinplugins.Source("jimeng")
require.NoError(t, err)
plugin, err := jsplugin.NewRegistry().RegisterFactory(source, jsplugin.Options{Key: "jimeng"})
require.NoError(t, err)
return plugin
}
func TestJimengSubmitStateDrivesQueryReqKey(t *testing.T) {
plugin := loadJimengPlugin(t)
submitValue, err := plugin.Engine.Call(t.Context(), "parseSubmitResponse", map[string]any{
"upstreamModel": "jimeng_vgfm_i2v_l20",
"requestBody": map[string]any{"images": []any{"https://cdn.example/frame.png"}},
}, map[string]any{"body": map[string]any{"code": 10000, "data": map[string]any{"task_id": "t1"}}})
require.NoError(t, err)
encoded, err := common.Marshal(submitValue)
require.NoError(t, err)
var submit map[string]any
require.NoError(t, common.Unmarshal(encoded, &submit))
state, ok := submit["state"].(map[string]any)
require.True(t, ok)
assert.Equal(t, "jimeng_vgfm_i2v_l20", state["req_key"])
queryValue, err := plugin.Engine.Call(t.Context(), "buildQueryRequest", map[string]any{
"taskId": "t1",
"action": "text_to_video",
"baseUrl": "https://jimeng.example",
"apiKey": "sk-test",
"state": map[string]any{"req_key": "custom_req_key"},
})
require.NoError(t, err)
queryEncoded, err := common.Marshal(queryValue)
require.NoError(t, err)
var query map[string]any
require.NoError(t, common.Unmarshal(queryEncoded, &query))
var body map[string]any
require.NoError(t, common.UnmarshalJsonStr(common.Interface2String(query["body"]), &body))
assert.Equal(t, "custom_req_key", body["req_key"])
assert.Equal(t, "t1", body["task_id"])
}
func TestJimengParseTaskResultUnknownStatus(t *testing.T) {
plugin := loadJimengPlugin(t)
value, err := plugin.Engine.Call(t.Context(), "parseTaskResult", map[string]any{}, map[string]any{
"code": 10000,
"data": map[string]any{"status": "weird"},
})
require.NoError(t, err)
encoded, err := common.Marshal(value)
require.NoError(t, err)
var result map[string]any
require.NoError(t, common.Unmarshal(encoded, &result))
assert.Equal(t, "UNKNOWN", result["status"])
assert.Contains(t, common.Interface2String(result["reason"]), "weird")
}
......@@ -7,7 +7,7 @@ export const meta = {
en: "Alibaba Cloud Bailian Wanxiang video generation (text-to-video and image-to-video)",
zh: "阿里云百炼万相视频生成(文生视频、图生视频)",
},
version: "1.0.0",
version: "1.0.1",
author: { name: "QuantumNous" },
channelTypes: [17],
models: [
......@@ -270,7 +270,7 @@ export function parseTaskResult(ctx, body) {
if (!reason) reason = "task failed";
return { status: "FAILURE", reason: reason };
}
return { status: "QUEUED" };
return { status: "UNKNOWN", reason: "unrecognized status: " + String(output.task_status || "") };
}
function artifactData(ctx) {
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "Volcengine Doubao Seedance video generation (text-to-video, image-to-video, and video-to-video)",
zh: "火山引擎豆包 Seedance 视频生成(文生视频、图生视频、视频生视频)",
},
version: "1.0.0",
version: "1.0.1",
author: { name: "QuantumNous" },
channelTypes: [54, 45], // VolcEngine-type channels serve Ark video models with the same wire format
models: [
......@@ -324,7 +324,7 @@ export function parseTaskResult(ctx, body) {
const reason = body.error && body.error.message ? body.error.message : body.status;
return { status: "FAILURE", progress: "100%", reason: reason };
}
return { status: "IN_PROGRESS", progress: "30%" };
return { status: "UNKNOWN", reason: "unrecognized status: " + String(body.status || "") };
}
function artifactData(ctx) {
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "Google Veo video generation on the Gemini API (text-to-video and image-to-video)",
zh: "Google Veo 视频生成(文生视频、图生视频),Gemini API 版本",
},
version: "1.0.0",
version: "1.0.1",
author: { name: "QuantumNous" },
channelTypes: [24],
models: ["veo-3.0-generate-001", "veo-3.0-fast-generate-001", "veo-3.1-generate-preview", "veo-3.1-fast-generate-preview"],
......@@ -209,7 +209,11 @@ export function buildQueryRequest(ctx) {
export function parseTaskResult(ctx, body) {
if (body.error && body.error.message) return { status: "FAILURE", progress: "100%", reason: body.error.message };
if (!body.done) return { status: "IN_PROGRESS", progress: "50%" };
// Google long-running operations omit `done` (proto3 default) while still
// running, so a missing key means in-progress; only a non-operation shape is
// unrecognized.
if (!body || typeof body !== "object" || !String(body.name || "").trim()) return { status: "UNKNOWN", reason: "unrecognized operation state" };
if (body.done !== true) return { status: "IN_PROGRESS", progress: "50%" };
const videos = ((body.response || {}).generateVideoResponse || {}).generatedVideos || [];
const uri = videos.length && videos[0].video ? videos[0].video.uri || "" : "";
return { taskId: utils.base64URL(body.name || ""), status: "SUCCESS", progress: "100%", remoteUrl: uri };
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "MiniMax Hailuo video generation (text-to-video, image-to-video, and MiniMax-H3 multimodal reference)",
zh: "MiniMax 海螺视频生成(文生视频、图生视频、MiniMax-H3 多模态参考生视频)",
},
version: "1.1.1",
version: "1.1.2",
author: { name: "QuantumNous" },
channelTypes: [35],
models: [
......@@ -465,8 +465,6 @@ export function buildQueryRequest(ctx) {
}
export function parseTaskResult(ctx, body) {
// The host calls this hook with an empty context, so the response envelope is
// the only way to tell a /v2 result from a /v1 one.
const apiError = h3APIError(body);
if (apiError) {
if (apiError.statusCode === 408 || apiError.statusCode === 429 || apiError.statusCode >= 500) throw new Error(apiError.message);
......@@ -475,7 +473,10 @@ export function parseTaskResult(ctx, body) {
const h3Task = h3QueryTask(body);
if (h3Task) {
const h3Statuses = { queued: "QUEUED", running: "IN_PROGRESS", succeeded: "SUCCESS", failed: "FAILURE", cancelled: "FAILURE" };
const h3Status = h3Statuses[h3Task.status] || "IN_PROGRESS";
const h3Status = h3Statuses[h3Task.status];
if (!h3Status) {
return { status: "UNKNOWN", reason: "unrecognized status: " + String(h3Task.status || "") };
}
const h3Result = { code: 0, status: h3Status, progress: h3Status === "QUEUED" ? "30%" : h3Status === "IN_PROGRESS" ? "50%" : "100%" };
if (h3Status === "SUCCESS") {
const url = trimmed(h3Task.content && h3Task.content.url);
......@@ -486,11 +487,17 @@ export function parseTaskResult(ctx, body) {
}
return h3Result;
}
if (body.base_resp && body.base_resp.status_code !== 0) {
return { code: body.base_resp.status_code || 0, status: "FAILURE", progress: "100%", reason: body.base_resp.status_msg || "" };
}
const base = body.base_resp || {};
const statuses = { Preparing: "IN_PROGRESS", Queueing: "IN_PROGRESS", Processing: "IN_PROGRESS", Success: "SUCCESS", Fail: "FAILURE" };
const status = statuses[body.status] || "IN_PROGRESS";
const status = statuses[body.status];
if (!status) {
return { status: "UNKNOWN", reason: "unrecognized status: " + String(body.status || "") };
}
const progress = status === "SUCCESS" || status === "FAILURE" ? "100%" : body.status === "Processing" ? "50%" : "30%";
const reason = base.status_code !== 0 ? base.status_msg || "" : status === "FAILURE" ? "task failed" : "";
const reason = status === "FAILURE" ? "task failed" : "";
return { code: base.status_code || 0, status: status, progress: progress, reason: reason };
}
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "Volcengine Jimeng video generation (text-to-video, image-to-video, and first-and-last-frame)",
zh: "火山引擎即梦视频生成(文生视频、图生视频、首尾帧)",
},
version: "1.0.0",
version: "1.0.1",
author: { name: "QuantumNous" },
channelTypes: [51],
models: ["jimeng_vgfm_t2v_l20"],
......@@ -270,10 +270,8 @@ function filePlaceholder(image) {
}
function queryReqKey(ctx) {
const data = (ctx && ctx.data) || {};
if (typeof data.req_key === "string" && data.req_key.trim()) return data.req_key.trim();
const req = (ctx && ctx.requestBody) || {};
if (typeof req.req_key === "string" && req.req_key.trim()) return req.req_key.trim();
const state = (ctx && ctx.state) || {};
if (typeof state.req_key === "string" && state.req_key.trim()) return state.req_key.trim();
if (ctx && ctx.action === "image_to_video") return "jimeng_vgfm_i2v_l20";
if (ctx && ctx.action === "first_tail_to_video") return "jimeng_i2v_first_tail_v30";
return "jimeng_vgfm_t2v_l20";
......@@ -349,7 +347,7 @@ export function parseSubmitResponse(ctx, resp) {
const body = resp.body || {};
if (body.code !== 10000) throw new Error(body.message || "jimeng submit failed");
if (!body.data || !body.data.task_id) throw new Error("missing task_id");
return { taskId: body.data.task_id, taskData: Object.assign({}, body, { req_key: submitReqKey(ctx) }) };
return { taskId: body.data.task_id, taskData: Object.assign({}, body, { req_key: submitReqKey(ctx) }), state: { req_key: submitReqKey(ctx) } };
}
export function extractUsage(ctx) {
......@@ -366,22 +364,20 @@ export function buildQueryRequest(ctx) {
export function parseTaskResult(ctx, body) {
const data = body.data || {};
let status = "";
let progress = "";
if (body.code !== 10000) {
status = "FAILURE";
progress = "100%";
return { code: body.code || 0, status: "FAILURE", progress: "100%", reason: body.message || "" };
}
if (data.status === "in_queue") {
status = "QUEUED";
progress = "10%";
} else if (data.status === "done") {
status = "SUCCESS";
progress = "100%";
const result = { code: 0, status: "QUEUED", progress: "10%", reason: "" };
if (data.video_url) result.url = data.video_url;
return result;
}
if (data.status === "done") {
const result = { code: 0, status: "SUCCESS", progress: "100%", reason: "" };
if (data.video_url) result.url = data.video_url;
return result;
}
const result = { code: body.code === 10000 ? 0 : body.code || 0, status: status, progress: progress, reason: body.code === 10000 ? "" : body.message || "" };
if (data.video_url) result.url = data.video_url;
return result;
return { code: 0, status: "UNKNOWN", reason: "unrecognized status: " + String(data.status || "") };
}
function artifactData(ctx) {
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "Kuaishou Kling video generation (text-to-video and image-to-video)",
zh: "快手可灵视频生成(文生视频、图生视频)",
},
version: "1.0.0",
version: "1.0.1",
author: { name: "QuantumNous" },
channelTypes: [50],
models: ["kling-v1", "kling-v1-6", "kling-v2-master"],
......@@ -286,7 +286,7 @@ export function parseTaskResult(ctx, body) {
const data = body.data || {};
const statuses = { submitted: "SUBMITTED", processing: "IN_PROGRESS", succeed: "SUCCESS", failed: "FAILURE" };
const status = statuses[data.task_status];
if (!status) throw new Error("unknown task status: " + data.task_status);
if (!status) return { status: "UNKNOWN", reason: "unknown task status: " + String(data.task_status || "") };
const videos = status === "SUCCESS" && data.task_result && data.task_result.videos ? data.task_result.videos : [];
const result = { code: body.code || 0, taskId: data.task_id, status: status, reason: data.task_status_msg || "" };
if (videos.length && videos[0].url) result.url = videos[0].url;
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "OpenAI Sora video generation (text-to-video, image-to-video, and remix)",
zh: "OpenAI Sora 视频生成(文生视频、图生视频、remix)",
},
version: "1.0.0",
version: "1.0.1",
channelTypes: [55, 1], // OpenAI-type channels natively serve sora with the same wire format
author: { name: "QuantumNous" },
models: ["sora-2", "sora-2-pro"],
......@@ -145,7 +145,9 @@ export function parseTaskResult(ctx, body) {
failed: "FAILURE",
cancelled: "FAILURE",
};
const result = { status: statuses[body.status] || "UNKNOWN" };
const mapped = statuses[body.status];
const result = { status: mapped || "UNKNOWN" };
if (!mapped) result.reason = "unrecognized status: " + String(body.status || "");
if (body.progress > 0 && body.progress < 100) result.progress = body.progress + "%";
if (result.status === "FAILURE") result.reason = body.error && body.error.message ? body.error.message : "task failed";
return result;
......
......@@ -9,7 +9,7 @@ export const meta = {
en: "SunoAPI project music and lyrics generation",
zh: "SunoAPI 项目 音乐与歌词生成",
},
version: "1.0.0",
version: "1.0.1",
author: { name: "QuantumNous" },
channelTypes: [36],
models: ["suno_music", "suno_lyrics"],
......@@ -133,19 +133,23 @@ export function extractUsage(ctx) {
return { clips: action === "lyrics" ? 1 : 2, action: action };
}
export function buildBatchQueryRequest(ctx, taskIds) {
export function buildBatchQueryRequest(ctx, tasks) {
return {
url: ctx.baseUrl + "/suno/fetch",
method: "POST",
headers: { "Content-Type": "application/json", Authorization: "Bearer " + ctx.apiKey },
body: { ids: taskIds },
body: {
ids: (tasks || []).map(function (task) {
return task.taskId;
}),
},
};
}
// Required v1 per-task hooks remain defined for contract compatibility. Suno's
// host polling path uses the batch hooks below.
export function buildQueryRequest(ctx) {
return buildBatchQueryRequest(ctx, (ctx.requestBody || {}).ids || []);
return buildBatchQueryRequest(ctx, [ctx]);
}
export function parseBatchResult(ctx, body) {
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "Google Veo video generation on Vertex AI (text-to-video and image-to-video)",
zh: "Google Veo 视频生成(文生视频、图生视频),Vertex AI 版本",
},
version: "1.0.0",
version: "1.0.1",
channelTypes: [41],
author: { name: "QuantumNous" },
models: ["veo-3.0-generate-001", "veo-3.0-fast-generate-001", "veo-3.1-generate-preview", "veo-3.1-fast-generate-preview"],
......@@ -225,7 +225,11 @@ export function buildQueryRequest(ctx) {
}
export function parseTaskResult(ctx, body) {
if (body.error && body.error.message) return { status: "FAILURE", progress: "100%", reason: body.error.message };
if (!body.done) return { status: "IN_PROGRESS", progress: "50%" };
// Google long-running operations omit `done` (proto3 default) while still
// running, so a missing key means in-progress; only a non-operation shape is
// unrecognized.
if (!body || typeof body !== "object" || !String(body.name || "").trim()) return { status: "UNKNOWN", reason: "unrecognized operation state" };
if (body.done !== true) return { status: "IN_PROGRESS", progress: "50%" };
const url = dataVideo(body.response || {});
return { status: "SUCCESS", progress: "100%", url: url, remoteUrl: url };
}
......
......@@ -7,7 +7,7 @@ export const meta = {
en: "Shengshu Vidu video generation (text-to-video, image-to-video, first-and-last-frame, and reference-to-video)",
zh: "生数 Vidu 视频生成(文生视频、图生视频、首尾帧、参考生视频)",
},
version: "1.0.0",
version: "1.0.1",
author: { name: "QuantumNous" },
channelTypes: [52],
models: ["viduq2", "viduq1", "vidu2.0", "vidu1.5"],
......@@ -265,7 +265,7 @@ export function buildQueryRequest(ctx) {
export function parseTaskResult(ctx, body) {
const statuses = { created: "SUBMITTED", queueing: "SUBMITTED", processing: "IN_PROGRESS", success: "SUCCESS", failed: "FAILURE" };
const status = statuses[body.state];
if (!status) throw new Error("unknown task state: " + body.state);
if (!status) return { status: "UNKNOWN", reason: "unknown task state: " + String(body.state || "") };
const url = body.creations && body.creations.length ? body.creations[0].url || "" : "";
const result = { status: status, reason: body.state === "failed" ? body.err_code || "" : "" };
if (url) result.url = url;
......
package plugins_test
import (
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/pkg/jsplugin"
builtinplugins "github.com/QuantumNous/new-api/plugins"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// Google long-running operations serialize proto3 defaults by omission: a
// still-running operation has no `done` key at all. Treating that as
// unrecognized would count every normal poll as a failure and fail the task
// at the poll-failure threshold while the video is still rendering.
func TestVeoParseTaskResultTreatsMissingDoneAsInProgress(t *testing.T) {
cases := []struct {
name string
body map[string]any
wantStatus string
}{
{
name: "running operation omits done",
body: map[string]any{"name": "operations/abc", "metadata": map[string]any{"@type": "x"}},
wantStatus: "IN_PROGRESS",
},
{
name: "explicit done false",
body: map[string]any{"name": "operations/abc", "done": false},
wantStatus: "IN_PROGRESS",
},
{
name: "body without operation name is unrecognized",
body: map[string]any{"foo": "bar"},
wantStatus: "UNKNOWN",
},
{
name: "operation error is failure",
body: map[string]any{"name": "operations/abc", "done": true, "error": map[string]any{"message": "quota exceeded"}},
wantStatus: "FAILURE",
},
}
for _, key := range []string{"google", "vertex-ai"} {
source, err := builtinplugins.Source(key)
require.NoError(t, err)
plugin, err := jsplugin.NewRegistry().RegisterFactory(source, jsplugin.Options{Key: key})
require.NoError(t, err)
for _, tc := range cases {
t.Run(key+"/"+tc.name, func(t *testing.T) {
value, err := plugin.Engine.Call(t.Context(), "parseTaskResult", map[string]any{}, tc.body)
require.NoError(t, err)
encoded, err := common.Marshal(value)
require.NoError(t, err)
var result map[string]any
require.NoError(t, common.Unmarshal(encoded, &result))
assert.Equal(t, tc.wantStatus, result["status"])
})
}
}
}
......@@ -76,8 +76,8 @@ type TaskAdaptor interface {
// ── Polling ──────────────────────────────────────────────────────
FetchTask(baseUrl, key string, body map[string]any, proxy string) (*http.Response, error)
ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, error)
FetchTask(baseUrl, key string, task *model.Task, proxy string) (*http.Response, error)
ParseTaskResult(task *model.Task, resp *http.Response, respBody []byte) (*relaycommon.TaskInfo, error)
}
// TaskSubmitResponse is the transport-independent result of parsing an
......@@ -87,6 +87,7 @@ type TaskSubmitResponse struct {
TaskData []byte
ClientResponse any
Immediate *relaycommon.TaskInfo
PluginState []byte
}
type OpenAIVideoConverter interface {
......
......@@ -22,12 +22,12 @@ import (
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
kitdto "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model"
pluginruntime "github.com/QuantumNous/new-api/pkg/jsplugin"
"github.com/QuantumNous/new-api/relay/channel"
relaycommon "github.com/QuantumNous/new-api/relay/common"
kitdto "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/service"
"github.com/gin-gonic/gin"
)
......@@ -56,6 +56,7 @@ type submitResponse struct {
TaskID string `json:"taskId"`
TaskData any `json:"taskData"`
Immediate *taskResult `json:"immediate"`
State any `json:"state"`
}
type taskResult struct {
Code int `json:"code"`
......@@ -67,12 +68,16 @@ type taskResult struct {
RemoteURL string `json:"remoteUrl"`
CompletionTokens float64 `json:"completionTokens"`
TotalTokens float64 `json:"totalTokens"`
State any `json:"state"`
}
var taskArtifactKeyPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._~-]{0,127}$`)
const maxTaskArtifacts = 64
// maxTaskPluginPersistedJSONBytes is the shared ceiling for taskData and plugin state.
const maxTaskPluginPersistedJSONBytes = 1 << 20
type TaskAdaptor struct {
plugin *pluginruntime.LoadedPlugin
info *relaycommon.RelayInfo
......@@ -497,10 +502,16 @@ func (a *TaskAdaptor) ParseResponse(c *gin.Context, resp *http.Response, info *r
immediate != nil,
time.Since(started).Milliseconds(),
)
pluginState, _ := encodeReturnedPluginState(value)
if len(pluginState) > maxTaskPluginPersistedJSONBytes {
logger.LogWarn(c, fmt.Sprintf("task plugin %s rejected oversized submit state (%d bytes)", a.plugin.Meta.Key, len(pluginState)))
pluginState = nil
}
return &channel.TaskSubmitResponse{
UpstreamTaskID: parsed.TaskID,
TaskData: taskData,
Immediate: immediate,
PluginState: pluginState,
}, nil
}
......@@ -508,50 +519,32 @@ func (a *TaskAdaptor) GetModelList() []string { return append([]string(nil), a.p
func (a *TaskAdaptor) GetChannelName() string { return a.plugin.Meta.Name }
func (a *TaskAdaptor) FetchMode() string { return a.plugin.Meta.FetchMode }
func (a *TaskAdaptor) FetchBatchTasks(baseURL, key string, taskIDs []string, proxy string) (*http.Response, error) {
ctx := map[string]any{"baseUrl": baseURL}
auth, err := resolveAuth(a.plugin.Meta.Auth, key, proxy)
func (a *TaskAdaptor) FetchBatchTasks(baseURL, key string, tasks []*model.Task, proxy string) (*http.Response, error) {
taskContexts := make([]map[string]any, 0, len(tasks))
for _, task := range tasks {
taskCtx, err := a.queryContext(task, key, baseURL, proxy)
if err != nil {
return nil, err
}
taskContexts = append(taskContexts, taskCtx)
}
ctx, err := a.batchQueryContext(key, baseURL, proxy, taskContexts)
if err != nil {
return nil, err
}
ctx["auth"] = auth
ctx["authHeader"] = auth["authHeader"]
if a.plugin.Meta.Auth.Type == "" || a.plugin.Meta.Auth.Type == "none" || a.plugin.Meta.Auth.Type == "api_key" {
ctx["apiKey"] = key
}
value, err := a.plugin.Engine.Call(context.Background(), "buildBatchQueryRequest", ctx, taskIDs)
value, err := a.plugin.Engine.Call(context.Background(), "buildBatchQueryRequest", ctx, taskContexts)
if err != nil {
return nil, err
}
return a.doFetchDescriptor(baseURL, proxy, value)
}
func (a *TaskAdaptor) FetchTask(baseURL, key string, body map[string]any, proxy string) (*http.Response, error) {
ctx := map[string]any{"taskId": body["task_id"], "action": body["action"], "requestBody": body, "baseUrl": baseURL}
// Query hooks are driver hooks and must see the same model identities as
// submit hooks. Polling has no relay info, so they arrive with the
// persisted task properties the caller puts in the fetch body.
originModel, _ := body["model"].(string)
upstreamModel, _ := body["upstream_model"].(string)
if upstreamModel == "" {
upstreamModel = originModel
}
ctx["model"] = originModel
ctx["upstreamModel"] = upstreamModel
auth, err := resolveAuth(a.plugin.Meta.Auth, key, proxy)
func (a *TaskAdaptor) FetchTask(baseURL, key string, task *model.Task, proxy string) (*http.Response, error) {
ctx, err := a.queryContext(task, key, baseURL, proxy)
if err != nil {
return nil, err
}
ctx["auth"] = auth
ctx["authHeader"] = auth["authHeader"]
if a.plugin.Meta.Auth.Type == "" || a.plugin.Meta.Auth.Type == "none" || a.plugin.Meta.Auth.Type == "api_key" {
ctx["apiKey"] = key
}
hook := "buildQueryRequest"
if a.plugin.Meta.FetchMode == "batch" && a.hasHook(context.Background(), "buildBatchQueryRequest") {
hook = "buildBatchQueryRequest"
}
value, err := a.plugin.Engine.Call(context.Background(), hook, ctx)
value, err := a.plugin.Engine.Call(context.Background(), "buildQueryRequest", ctx)
if err != nil {
return nil, err
}
......@@ -616,14 +609,27 @@ func (a *TaskAdaptor) doFetchDescriptor(baseURL, proxy string, value any) (*http
return resp, nil
}
func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTaskResult, error) {
func (a *TaskAdaptor) ParseBatchResult(tasks []*model.Task, resp *http.Response, body []byte) (map[string]*service.BatchTaskResult, error) {
started := time.Now()
input := any(string(body))
var decoded any
if common.Unmarshal(body, &decoded) == nil {
input = decoded
}
value, err := a.plugin.Engine.Call(context.Background(), "parseBatchResult", map[string]any{}, input)
key, baseURL, proxy := a.queryCredentials()
taskContexts := make([]map[string]any, 0, len(tasks))
for _, task := range tasks {
taskCtx, err := a.queryContext(task, key, baseURL, proxy)
if err != nil {
return nil, err
}
taskContexts = append(taskContexts, taskCtx)
}
ctx, err := a.batchQueryContext(key, baseURL, proxy, taskContexts)
if err != nil {
return nil, err
}
value, err := a.plugin.Engine.Call(context.Background(), "parseBatchResult", ctx, input, hookHTTPResponse(resp))
if err != nil {
logger.LogDebug(context.Background(), "task_plugin subsystem=adaptor event=parse_batch_failed plugin=%q reason=hook_failed body_bytes=%d elapsed_ms=%d", a.plugin.Meta.Key, len(body), time.Since(started).Milliseconds())
return nil, err
......@@ -639,6 +645,7 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa
StartTime int64 `json:"startTime"`
FinishTime int64 `json:"finishTime"`
Data any `json:"data"`
State any `json:"state"`
}
if err = convert(value, &parsed); err != nil {
logger.LogDebug(context.Background(), "task_plugin subsystem=adaptor event=parse_batch_failed plugin=%q reason=invalid_result body_bytes=%d elapsed_ms=%d", a.plugin.Meta.Key, len(body), time.Since(started).Milliseconds())
......@@ -651,12 +658,27 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa
continue
}
info := relaycommon.TaskInfo{TaskID: item.TaskID, Status: item.Status, Progress: item.Progress, Reason: item.Reason, Url: item.URL}
if item.State != nil {
pluginState, marshalErr := common.Marshal(item.State)
if marshalErr != nil || len(pluginState) > maxTaskPluginPersistedJSONBytes {
logger.LogWarn(context.Background(), fmt.Sprintf("task plugin %s rejected invalid or oversized poll state", a.plugin.Meta.Key))
} else {
info.PluginState = pluginState
}
}
if hasCompletionUsage {
usageBody := item.Data
if usageBody == nil {
usageBody = jsonValue(item)
}
facts, hookErr := a.plugin.Engine.Call(context.Background(), "extractUsageOnComplete", nil, jsonValue(&info), usageBody)
itemCtx := ctx
for _, taskCtx := range taskContexts {
if fmt.Sprint(taskCtx["taskId"]) == item.TaskID {
itemCtx = taskCtx
break
}
}
facts, hookErr := a.plugin.Engine.Call(context.Background(), "extractUsageOnComplete", itemCtx, jsonValue(&info), usageBody)
if hookErr == nil {
a.applyCompletionUsageFacts(&info, facts)
}
......@@ -675,14 +697,22 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa
return results, nil
}
func (a *TaskAdaptor) ParseTaskResult(body []byte) (*relaycommon.TaskInfo, error) {
func (a *TaskAdaptor) ParseTaskResult(task *model.Task, resp *http.Response, body []byte) (*relaycommon.TaskInfo, error) {
started := time.Now()
input := any(string(body))
var decoded any
if common.Unmarshal(body, &decoded) == nil {
input = decoded
}
value, err := a.plugin.Engine.Call(context.Background(), "parseTaskResult", map[string]any{}, input)
key, baseURL, proxy := a.queryCredentials()
if task != nil && task.PrivateData.Key != "" {
key = task.PrivateData.Key
}
ctx, err := a.queryContext(task, key, baseURL, proxy)
if err != nil {
return nil, err
}
value, err := a.plugin.Engine.Call(context.Background(), "parseTaskResult", ctx, input, hookHTTPResponse(resp))
if err != nil {
logger.LogDebug(context.Background(), "task_plugin subsystem=adaptor event=parse_task_failed plugin=%q reason=hook_failed body_bytes=%d elapsed_ms=%d", a.plugin.Meta.Key, len(body), time.Since(started).Milliseconds())
return nil, err
......@@ -703,10 +733,17 @@ func (a *TaskAdaptor) ParseTaskResult(body []byte) (*relaycommon.TaskInfo, error
CompletionTokens: positiveInt(parsed.CompletionTokens),
TotalTokens: positiveInt(parsed.TotalTokens),
}
if pluginState, present := encodeReturnedPluginState(value); present {
if len(pluginState) > maxTaskPluginPersistedJSONBytes {
logger.LogWarn(context.Background(), fmt.Sprintf("task plugin %s rejected oversized poll state (%d bytes)", a.plugin.Meta.Key, len(pluginState)))
} else {
result.PluginState = pluginState
}
}
// The raw polling response only exists at this boundary. Capture upstream
// units here so the host settlement path can consume them from TaskInfo.
if a.hasHook(context.Background(), "extractUsageOnComplete") {
facts, hookErr := a.plugin.Engine.Call(context.Background(), "extractUsageOnComplete", nil, jsonValue(result), input)
facts, hookErr := a.plugin.Engine.Call(context.Background(), "extractUsageOnComplete", ctx, jsonValue(result), input)
if hookErr == nil {
a.applyCompletionUsageFacts(result, facts)
}
......@@ -901,15 +938,125 @@ func taskArtifactContext(task *model.Task) (map[string]any, error) {
if task.PrivateData.Execution != nil && task.PrivateData.Execution.TaskPlugin != nil {
producerVersion = task.PrivateData.Execution.TaskPlugin.Version
}
var state any
if len(task.PrivateData.PluginState) > 0 {
if err := common.Unmarshal(task.PrivateData.PluginState, &state); err != nil {
return nil, fmt.Errorf("plugin state is invalid")
}
}
return map[string]any{
"taskId": task.TaskID,
"status": string(task.Status),
"action": task.Action,
"data": data,
"state": state,
"producerVersion": producerVersion,
}, nil
}
func (a *TaskAdaptor) queryContext(task *model.Task, key, baseURL, proxy string) (map[string]any, error) {
ctx := map[string]any{
"taskId": "",
"publicTaskId": "",
"action": "",
"model": "",
"upstreamModel": "",
"baseUrl": baseURL,
"data": nil,
"state": nil,
}
if task != nil {
originModel := task.Properties.OriginModelName
upstreamModel := task.Properties.UpstreamModelName
if upstreamModel == "" {
upstreamModel = originModel
}
ctx["taskId"] = task.GetUpstreamTaskID()
ctx["publicTaskId"] = task.TaskID
ctx["action"] = constant.NormalizeTaskAction(task.Action)
ctx["model"] = originModel
ctx["upstreamModel"] = upstreamModel
if len(task.Data) > 0 {
var data any
if err := common.Unmarshal(task.Data, &data); err != nil {
return nil, fmt.Errorf("task data is invalid")
}
ctx["data"] = data
}
if len(task.PrivateData.PluginState) > 0 {
var state any
if err := common.Unmarshal(task.PrivateData.PluginState, &state); err != nil {
return nil, fmt.Errorf("plugin state is invalid")
}
ctx["state"] = state
}
if task.PrivateData.Key != "" {
key = task.PrivateData.Key
}
}
auth, err := resolveAuth(a.plugin.Meta.Auth, key, proxy)
if err != nil {
return nil, err
}
ctx["auth"] = auth
ctx["authHeader"] = auth["authHeader"]
if a.plugin.Meta.Auth.Type == "" || a.plugin.Meta.Auth.Type == "none" || a.plugin.Meta.Auth.Type == "api_key" {
ctx["apiKey"] = key
}
return ctx, nil
}
func (a *TaskAdaptor) batchQueryContext(key, baseURL, proxy string, tasks []map[string]any) (map[string]any, error) {
ctx := map[string]any{"baseUrl": baseURL, "tasks": tasks}
auth, err := resolveAuth(a.plugin.Meta.Auth, key, proxy)
if err != nil {
return nil, err
}
ctx["auth"] = auth
ctx["authHeader"] = auth["authHeader"]
if a.plugin.Meta.Auth.Type == "" || a.plugin.Meta.Auth.Type == "none" || a.plugin.Meta.Auth.Type == "api_key" {
ctx["apiKey"] = key
}
return ctx, nil
}
func (a *TaskAdaptor) queryCredentials() (key, baseURL, proxy string) {
if a.info == nil || !a.info.HasChannelMeta() {
return "", "", ""
}
return a.info.ApiKey, a.info.ChannelBaseUrl, a.info.ChannelSetting.Proxy
}
func hookHTTPResponse(resp *http.Response) map[string]any {
headers := map[string]string{}
status := 0
if resp != nil {
status = resp.StatusCode
for name, values := range resp.Header {
if len(values) > 0 {
headers[name] = values[0]
}
}
}
return map[string]any{"status": status, "headers": headers}
}
func encodeReturnedPluginState(value any) ([]byte, bool) {
object, ok := value.(map[string]any)
if !ok {
return nil, false
}
state, exists := object["state"]
if !exists || state == nil {
return nil, false
}
data, err := common.Marshal(state)
if err != nil {
return nil, false
}
return data, true
}
func validateTaskArtifacts(value any) ([]channel.TaskArtifact, error) {
encoded, err := common.Marshal(value)
if err != nil {
......
......@@ -503,12 +503,15 @@ func TestTaskAdaptorMapsJSContract(t *testing.T) {
assert.Empty(t, recorder.Body.String(), "response parsing must not write before the durable task barrier")
assert.Equal(t, map[string]float64{"seconds": 7}, adaptor.AdjustBillingOnSubmit(info, []byte(`{"seconds":7}`)))
queryResp, err := adaptor.FetchTask(server.URL, "secret", map[string]any{"task_id": parsed.UpstreamTaskID, "action": info.Action}, "")
queryResp, err := adaptor.FetchTask(server.URL, "secret", &model.Task{
Action: info.Action,
PrivateData: model.TaskPrivateData{UpstreamTaskID: parsed.UpstreamTaskID},
}, "")
require.NoError(t, err)
queryBody, err := io.ReadAll(queryResp.Body)
require.NoError(t, err)
require.NoError(t, queryResp.Body.Close())
result, err := adaptor.ParseTaskResult(queryBody)
result, err := adaptor.ParseTaskResult(&model.Task{}, queryResp, queryBody)
require.NoError(t, err)
assert.Equal(t, "SUCCESS", result.Status)
assert.Equal(t, "https://cdn.example/video.mp4", result.Url)
......@@ -866,7 +869,7 @@ export function extractUsageOnComplete(task, result, body) { return (body || {})
adaptor, _, _ := newRequest(t, map[string]any{})
body, marshalErr := common.Marshal(map[string]any{"completionUsage": testCase.usage})
require.NoError(t, marshalErr)
result, parseErr := adaptor.ParseTaskResult(body)
result, parseErr := adaptor.ParseTaskResult(&model.Task{}, &http.Response{StatusCode: http.StatusOK, Header: make(http.Header)}, body)
require.NoError(t, parseErr)
assert.Nil(t, result.UsageFacts)
assert.Zero(t, result.TotalTokens)
......@@ -877,7 +880,7 @@ export function extractUsageOnComplete(task, result, body) { return (body || {})
adaptor, _, _ := newRequest(t, map[string]any{})
body, err := common.Marshal(map[string]any{"completionUsage": map[string]any{"tokens": 500000}})
require.NoError(t, err)
result, err := adaptor.ParseTaskResult(body)
result, err := adaptor.ParseTaskResult(&model.Task{}, &http.Response{StatusCode: http.StatusOK, Header: make(http.Header)}, body)
require.NoError(t, err)
assert.EqualValues(t, 500000, result.UsageFacts["tokens"])
})
......@@ -916,7 +919,7 @@ export function extractUsageOnComplete() { return {units: 3.5}; }
body, err := common.Marshal(map[string]any{})
require.NoError(t, err)
result, err := adaptor.ParseTaskResult(body)
result, err := adaptor.ParseTaskResult(&model.Task{}, &http.Response{StatusCode: http.StatusOK, Header: make(http.Header)}, body)
require.NoError(t, err)
assert.Equal(t, 3.5, result.UsageFacts["units"])
})
......@@ -925,7 +928,7 @@ export function extractUsageOnComplete() { return {units: 3.5}; }
adaptor, _, _ := newRequest(t, map[string]any{})
body, err := common.Marshal(map[string]any{"completionUsage": map[string]any{"upstreamUnits": 5000}})
require.NoError(t, err)
result, err := adaptor.ParseTaskResult(body)
result, err := adaptor.ParseTaskResult(&model.Task{}, &http.Response{StatusCode: http.StatusOK, Header: make(http.Header)}, body)
require.NoError(t, err)
assert.Equal(t, 5000, result.TotalTokens)
assert.EqualValues(t, 5000, result.UsageFacts["upstreamUnits"])
......@@ -988,7 +991,7 @@ export function parseTaskResult(ctx, body) { return {status: "SUCCESS", completi
require.NoError(t, err)
adaptor := New(plugin)
result, err := adaptor.ParseTaskResult([]byte(`{"completion":13,"total":17}`))
result, err := adaptor.ParseTaskResult(&model.Task{}, &http.Response{StatusCode: http.StatusOK, Header: make(http.Header)}, []byte(`{"completion":13,"total":17}`))
require.NoError(t, err)
assert.Equal(t, 13, result.CompletionTokens)
assert.Equal(t, 17, result.TotalTokens)
......@@ -1093,7 +1096,7 @@ export function buildSubmitRequest(ctx) { return { url: ctx.baseUrl + "/submit",
export function parseSubmitResponse(ctx, resp) { return { taskId: resp.body.id }; }
export function buildQueryRequest(ctx) { return { url: ctx.baseUrl + "/tasks/" + ctx.taskId }; }
export function parseTaskResult(ctx, body) { return { taskId: body.id, status: "SUCCESS" }; }
export function buildBatchQueryRequest(ctx, taskIds) { return { url: ctx.baseUrl + "/batch", method: "POST", headers: { "X-Plugin": "batch" }, body: { ids: taskIds } }; }
export function buildBatchQueryRequest(ctx, tasks) { return { url: ctx.baseUrl + "/batch", method: "POST", headers: { "X-Plugin": "batch" }, body: { ids: (tasks || []).map(function (task) { return task.taskId; }) } }; }
export function parseBatchResult(ctx, body) {
return body.items.map(function (item) {
return { taskId: item.id, action: item.action, status: item.status, progress: item.progress, url: (item.urls || [])[0] || "", finishTime: item.finish || 0, data: item };
......@@ -1128,13 +1131,17 @@ func TestTaskAdaptorBatchBridge(t *testing.T) {
adaptor := New(plugin)
require.Equal(t, "batch", adaptor.FetchMode())
resp, err := adaptor.FetchBatchTasks(server.URL, "secret", []string{"task-a", "task-b"}, "")
tasks := []*model.Task{
{PrivateData: model.TaskPrivateData{UpstreamTaskID: "task-a"}},
{PrivateData: model.TaskPrivateData{UpstreamTaskID: "task-b"}},
}
resp, err := adaptor.FetchBatchTasks(server.URL, "secret", tasks, "")
require.NoError(t, err)
defer resp.Body.Close()
payload, err := io.ReadAll(resp.Body)
require.NoError(t, err)
results, err := adaptor.ParseBatchResult(payload)
results, err := adaptor.ParseBatchResult(tasks, resp, payload)
require.NoError(t, err)
require.Len(t, results, 2, "entry without taskId must be skipped")
......@@ -1233,26 +1240,176 @@ export function parseTaskResult(){return {status:"SUCCESS"}}
testCases := []struct {
name string
body map[string]any
task *model.Task
want string
}{
{
name: "mapped model",
body: map[string]any{"task_id": "t1", "model": "alias", "upstream_model": "declared-model"},
task: &model.Task{
Properties: model.Properties{OriginModelName: "alias", UpstreamModelName: "declared-model"},
PrivateData: model.TaskPrivateData{UpstreamTaskID: "t1"},
},
want: "/tasks/alias/declared-model/t1",
},
{
name: "unmapped model falls back to the origin name",
body: map[string]any{"task_id": "t1", "model": "alias"},
task: &model.Task{
Properties: model.Properties{OriginModelName: "alias"},
PrivateData: model.TaskPrivateData{UpstreamTaskID: "t1"},
},
want: "/tasks/alias/alias/t1",
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
resp, fetchErr := adaptor.FetchTask(server.URL, "secret", testCase.body, "")
resp, fetchErr := adaptor.FetchTask(server.URL, "secret", testCase.task, "")
require.NoError(t, fetchErr)
require.NoError(t, resp.Body.Close())
assert.Equal(t, testCase.want, requested)
})
}
}
func TestTaskAdaptorQueryContextOmitsRequestBody(t *testing.T) {
service.InitHttpClient()
var captured map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
require.NoError(t, common.DecodeJson(r.Body, &captured))
_, _ = w.Write([]byte(`{"ok":true}`))
}))
defer server.Close()
source := `
export const meta = {apiVersion:1,key:"query-ctx",name:"Query Ctx",version:"1.0.0",author:{name:"Test"},models:["alias"],fetchMode:"per_task"};
export function buildSubmitRequest(ctx){return {url:ctx.baseUrl+"/submit"}}
export function parseSubmitResponse(){return {taskId:"1",state:{req_key:"from-submit"}}}
export function buildQueryRequest(ctx){
return {url:ctx.baseUrl+"/query",method:"POST",body:{
keys: Object.keys(ctx).sort(),
taskId: ctx.taskId,
publicTaskId: ctx.publicTaskId,
action: ctx.action,
model: ctx.model,
upstreamModel: ctx.upstreamModel,
data: ctx.data,
state: ctx.state,
hasRequestBody: Object.prototype.hasOwnProperty.call(ctx, "requestBody")
}};
}
export function parseTaskResult(ctx, body, response){
return {status:"IN_PROGRESS",reason:String(response && response.status),url:ctx.taskId,state:{round:"poll"}};
}
`
plugin, err := pluginruntime.NewRegistry().Register(source, pluginruntime.Options{})
require.NoError(t, err)
adaptor := New(plugin)
adaptor.Init(&relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{ChannelBaseUrl: server.URL, ApiKey: "secret"}})
task := &model.Task{
TaskID: "task_public",
Action: constant.TaskActionImageToVideo,
Properties: model.Properties{
OriginModelName: "alias",
UpstreamModelName: "declared",
},
Data: []byte(`{"snapshot":true}`),
PrivateData: model.TaskPrivateData{
UpstreamTaskID: "upstream-1",
PluginState: []byte(`{"req_key":"kept"}`),
},
}
resp, err := adaptor.FetchTask(server.URL, "secret", task, "")
require.NoError(t, err)
require.NoError(t, resp.Body.Close())
assert.Equal(t, "upstream-1", captured["taskId"])
assert.Equal(t, "task_public", captured["publicTaskId"])
assert.Equal(t, constant.NormalizeTaskAction(constant.TaskActionImageToVideo), captured["action"])
assert.Equal(t, "alias", captured["model"])
assert.Equal(t, "declared", captured["upstreamModel"])
assert.Equal(t, map[string]any{"snapshot": true}, captured["data"])
assert.Equal(t, map[string]any{"req_key": "kept"}, captured["state"])
assert.Equal(t, false, captured["hasRequestBody"])
keys, ok := captured["keys"].([]any)
require.True(t, ok)
assert.NotContains(t, keys, "requestBody")
result, err := adaptor.ParseTaskResult(task, &http.Response{StatusCode: http.StatusTeapot, Header: make(http.Header)}, []byte(`{"ok":true}`))
require.NoError(t, err)
assert.Equal(t, "IN_PROGRESS", result.Status)
assert.Equal(t, "418", result.Reason)
assert.Equal(t, "upstream-1", result.Url)
assert.JSONEq(t, `{"round":"poll"}`, string(result.PluginState))
}
func TestTaskAdaptorParseSubmitResponsePersistsState(t *testing.T) {
service.InitHttpClient()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte(`{"id":"upstream-1"}`))
}))
defer server.Close()
source := `
export const meta = {apiVersion:1,key:"submit-state",name:"Submit State",version:"1.0.0",author:{name:"Test"},models:["m"],fetchMode:"per_task"};
export function buildSubmitRequest(ctx){return {url:ctx.baseUrl+"/submit",method:"POST",body:{}}}
export function parseSubmitResponse(){return {taskId:"upstream-1",state:{req_key:"from-submit"}}}
export function buildQueryRequest(ctx){return {url:ctx.baseUrl+"/query"}}
export function parseTaskResult(){return {status:"SUCCESS"}}
`
plugin, err := pluginruntime.NewRegistry().Register(source, pluginruntime.Options{})
require.NoError(t, err)
adaptor := New(plugin)
info := &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{ChannelBaseUrl: server.URL, ApiKey: "secret"}, TaskRelayInfo: &relaycommon.TaskRelayInfo{}}
adaptor.Init(info)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest(http.MethodPost, "/v1/videos", nil)
c.Set("task_request", relaycommon.TaskSubmitReq{Prompt: "hello"})
require.Nil(t, adaptor.ValidateRequestAndSetAction(c, info))
body, err := adaptor.BuildRequestBody(c, info)
require.NoError(t, err)
resp, err := adaptor.DoRequest(c, info, body)
require.NoError(t, err)
parsed, taskErr := adaptor.ParseResponse(c, resp, info)
require.Nil(t, taskErr)
require.NotNil(t, parsed)
assert.JSONEq(t, `{"req_key":"from-submit"}`, string(parsed.PluginState))
}
func TestTaskAdaptorBatchQueryReceivesTaskObjects(t *testing.T) {
service.InitHttpClient()
var captured map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
require.NoError(t, common.DecodeJson(r.Body, &captured))
_, _ = w.Write([]byte(`{"items":[]}`))
}))
defer server.Close()
source := `
export const meta = {apiVersion:1,key:"batch-ctx",name:"Batch Ctx",version:"1.0.0",author:{name:"Test"},models:["m"],fetchMode:"batch"};
export function buildSubmitRequest(ctx){return {url:ctx.baseUrl+"/submit"}}
export function parseSubmitResponse(){return {taskId:"1"}}
export function buildQueryRequest(ctx){return {url:ctx.baseUrl+"/q"}}
export function parseTaskResult(){return {status:"SUCCESS"}}
export function buildBatchQueryRequest(ctx, tasks){
return {url:ctx.baseUrl+"/batch",method:"POST",body:{
ids: (tasks||[]).map(function(task){return task.taskId;}),
models: (tasks||[]).map(function(task){return task.model;}),
hasRequestBody: (tasks||[]).some(function(task){return Object.prototype.hasOwnProperty.call(task,"requestBody");})
}};
}
export function parseBatchResult(){return [];}
`
plugin, err := pluginruntime.NewRegistry().Register(source, pluginruntime.Options{})
require.NoError(t, err)
adaptor := New(plugin)
tasks := []*model.Task{
{Properties: model.Properties{OriginModelName: "model-a"}, PrivateData: model.TaskPrivateData{UpstreamTaskID: "task-a"}},
{Properties: model.Properties{OriginModelName: "model-b"}, PrivateData: model.TaskPrivateData{UpstreamTaskID: "task-b"}},
}
resp, err := adaptor.FetchBatchTasks(server.URL, "secret", tasks, "")
require.NoError(t, err)
require.NoError(t, resp.Body.Close())
assert.Equal(t, []any{"task-a", "task-b"}, captured["ids"])
assert.Equal(t, []any{"model-a", "model-b"}, captured["models"])
assert.Equal(t, false, captured["hasRequestBody"])
}
......@@ -119,19 +119,19 @@ type RelayInfo struct {
ReasoningEffort string
// ReasoningConversion is the suffix-derived reasoning intent attached
// after model mapping. Converters read it via ReasoningState().
ReasoningConversion *dto.ReasoningConversionState
UserSetting dto.UserSetting
UserEmail string
UserQuota int
RelayFormat types.RelayFormat
SendResponseCount int
ReasoningConversion *dto.ReasoningConversionState
UserSetting dto.UserSetting
UserEmail string
UserQuota int
RelayFormat types.RelayFormat
SendResponseCount int
// ClaudeToChatStreamState / ChatToGeminiStreamState hold per-attempt
// stream converters. InitChannelMeta nils them so a retry cannot resume a
// dirty converter (advanced tool index / finalized).
ClaudeToChatStreamState any
ChatToGeminiStreamState any
ReceivedResponseCount int
FinalPreConsumedQuota int // 最终预消耗的配额
FinalPreConsumedQuota int // 最终预消耗的配额
// ForcePreConsume 为 true 时禁用 BillingSession 的信任额度旁路,
// 强制预扣全额。用于异步任务(视频/音乐生成等),因为请求返回后任务仍在运行,
// 必须在提交前锁定全额。
......@@ -1018,16 +1018,17 @@ func (t *TaskSubmitReq) UnmarshalMetadata(v any) error {
}
type TaskInfo struct {
Code int `json:"code"`
TaskID string `json:"task_id"`
Status string `json:"status"`
Reason string `json:"reason,omitempty"`
Url string `json:"url,omitempty"`
RemoteUrl string `json:"remote_url,omitempty"`
Progress string `json:"progress,omitempty"`
CompletionTokens int `json:"completion_tokens,omitempty"` // 用于按倍率计费
TotalTokens int `json:"total_tokens,omitempty"` // 用于按倍率计费
UsageFacts map[string]any `json:"usage_facts,omitempty"`
Code int `json:"code"`
TaskID string `json:"task_id"`
Status string `json:"status"`
Reason string `json:"reason,omitempty"`
Url string `json:"url,omitempty"`
RemoteUrl string `json:"remote_url,omitempty"`
Progress string `json:"progress,omitempty"`
CompletionTokens int `json:"completion_tokens,omitempty"` // 用于按倍率计费
TotalTokens int `json:"total_tokens,omitempty"` // 用于按倍率计费
UsageFacts map[string]any `json:"usage_facts,omitempty"`
PluginState json.RawMessage `json:"plugin_state,omitempty"`
}
func FailTaskInfo(reason string) *TaskInfo {
......
......@@ -32,6 +32,7 @@ type TaskSubmitResult struct {
Platform constant.TaskPlatform
Quota int
Immediate *relaycommon.TaskInfo
PluginState []byte
//PerCallPrice types.PriceData
}
......@@ -381,6 +382,7 @@ func RelayTaskSubmit(c *gin.Context, info *relaycommon.RelayInfo) (*TaskSubmitRe
Platform: platform,
Quota: finalQuota,
Immediate: parsed.Immediate,
PluginState: parsed.PluginState,
}, nil
}
......@@ -517,12 +519,7 @@ func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte {
return nil
}
resp, err := adaptor.FetchTask(baseURL, channelModel.Key, map[string]any{
"task_id": task.GetUpstreamTaskID(),
"action": constant.NormalizeTaskAction(task.Action),
"model": task.Properties.OriginModelName,
"upstream_model": task.Properties.UpstreamModelName,
}, proxy)
resp, err := adaptor.FetchTask(baseURL, channelModel.Key, task, proxy)
if err != nil || resp == nil {
return nil
}
......@@ -532,7 +529,7 @@ func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte {
return nil
}
ti, err := adaptor.ParseTaskResult(body)
ti, err := adaptor.ParseTaskResult(task, resp, body)
if err != nil || ti == nil {
return nil
}
......
......@@ -1343,10 +1343,12 @@ type mockAdaptor struct {
}
func (m *mockAdaptor) Init(_ *relaycommon.RelayInfo) {}
func (m *mockAdaptor) FetchTask(string, string, map[string]any, string) (*http.Response, error) {
func (m *mockAdaptor) FetchTask(string, string, *model.Task, string) (*http.Response, error) {
return nil, nil
}
func (m *mockAdaptor) ParseTaskResult(*model.Task, *http.Response, []byte) (*relaycommon.TaskInfo, error) {
return nil, nil
}
func (m *mockAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) { return nil, nil }
func (m *mockAdaptor) AdjustBillingOnComplete(_ *model.Task, _ *relaycommon.TaskInfo) int {
return m.adjustReturn
}
......
......@@ -62,3 +62,24 @@ func TestBuildTaskPluginViewRewritesOnlyStructuredTaskIDFields(t *testing.T) {
assert.Equal(t, privateTaskID, nested[1])
}
func TestBuildTaskPluginViewOmitsPrivatePollState(t *testing.T) {
task := &model.Task{
TaskID: "task_public_view",
Data: []byte(`{"ok":true}`),
PrivateData: model.TaskPrivateData{
PluginState: []byte(`{"req_key":"secret"}`),
PollFailures: 7,
},
}
view, err := BuildTaskPluginView(task)
require.NoError(t, err)
encoded, err := common.Marshal(view)
require.NoError(t, err)
var payload map[string]any
require.NoError(t, common.Unmarshal(encoded, &payload))
assert.NotContains(t, payload, "plugin_state")
assert.NotContains(t, payload, "poll_failures")
assert.NotContains(t, payload, "private_data")
}
......@@ -18,7 +18,6 @@ import (
"github.com/QuantumNous/new-api/pkg/billingexpr"
"github.com/QuantumNous/new-api/relay/channel/task/taskcommon"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/bytedance/gopkg/util/gopool"
"github.com/samber/lo"
......@@ -27,8 +26,8 @@ import (
// TaskPollingAdaptor 定义轮询所需的最小适配器接口,避免 service -> relay 的循环依赖
type TaskPollingAdaptor interface {
Init(info *relaycommon.RelayInfo)
FetchTask(baseURL string, key string, body map[string]any, proxy string) (*http.Response, error)
ParseTaskResult(body []byte) (*relaycommon.TaskInfo, error)
FetchTask(baseURL string, key string, task *model.Task, proxy string) (*http.Response, error)
ParseTaskResult(task *model.Task, resp *http.Response, body []byte) (*relaycommon.TaskInfo, error)
// AdjustBillingOnComplete 在任务到达终态(成功/失败)时由轮询循环调用。
// 返回正数触发差额结算(补扣/退还),返回 0 保持预扣费金额不变。
AdjustBillingOnComplete(task *model.Task, taskResult *relaycommon.TaskInfo) int
......@@ -37,10 +36,21 @@ type TaskPollingAdaptor interface {
type BatchTaskPollingAdaptor interface {
TaskPollingAdaptor
FetchMode() string
FetchBatchTasks(baseURL, key string, taskIDs []string, proxy string) (*http.Response, error)
ParseBatchResult(body []byte) (map[string]*BatchTaskResult, error)
FetchBatchTasks(baseURL, key string, tasks []*model.Task, proxy string) (*http.Response, error)
ParseBatchResult(tasks []*model.Task, resp *http.Response, body []byte) (map[string]*BatchTaskResult, error)
}
const (
pollClassOK = "ok"
pollClassOtherClient = "other_client"
pollClassNotFound = "not_found"
pollClassAuth = "auth"
pollClassTransient = "transient"
pollClassUnrecognized = "unrecognized"
pollClassHookError = "hook_error"
pollClassTransport = "transport_error"
)
type BatchTaskResult struct {
TaskInfo relaycommon.TaskInfo
Action string
......@@ -258,24 +268,39 @@ func updateBatchTasks(ctx context.Context, adaptor BatchTaskPollingAdaptor, chan
if baseURL == "" {
baseURL = constant.GetChannelBaseURL(ch.Type)
}
resp, err := adaptor.FetchBatchTasks(baseURL, ch.Key, taskIds, proxy)
tasks := make([]*model.Task, 0, len(taskIds))
for _, upstreamID := range taskIds {
if task := taskM[upstreamID]; task != nil {
tasks = append(tasks, task)
}
}
info := &relaycommon.RelayInfo{}
info.ChannelMeta = &relaycommon.ChannelMeta{ChannelBaseUrl: baseURL}
info.ApiKey = ch.Key
adaptor.Init(info)
resp, err := adaptor.FetchBatchTasks(baseURL, ch.Key, tasks, proxy)
if err != nil {
common.SysLog(fmt.Sprintf("Get Task Do req error: %v", err))
return err
}
if resp.StatusCode != http.StatusOK {
logger.LogError(ctx, fmt.Sprintf("Get Task status code: %d", resp.StatusCode))
return fmt.Errorf("Get Task status code: %d", resp.StatusCode)
return recordPollFailureForTasks(ctx, adaptor, tasks, pollClassTransport, 0, err.Error())
}
defer resp.Body.Close()
responseBody, err := io.ReadAll(resp.Body)
if err != nil {
common.SysLog(fmt.Sprintf("Get Suno Task parse body error: %v", err))
return err
}
responseItems, err := adaptor.ParseBatchResult(responseBody)
return recordPollFailureForTasks(ctx, adaptor, tasks, pollClassTransport, resp.StatusCode, err.Error())
}
switch classifyPollHTTP(resp.StatusCode) {
case pollClassNotFound:
return failTasksFromPoll(ctx, adaptor, tasks, fmt.Sprintf("upstream task not found (HTTP %d)", resp.StatusCode))
case pollClassAuth:
logger.LogWarn(ctx, fmt.Sprintf("task poll auth failure channel_id=%d http=%d", channelId, resp.StatusCode))
return recordPollFailureForTasks(ctx, adaptor, tasks, pollClassAuth, resp.StatusCode, "")
case pollClassTransient:
return recordPollFailureForTasks(ctx, adaptor, tasks, pollClassTransient, resp.StatusCode, "")
}
responseItems, err := adaptor.ParseBatchResult(tasks, resp, responseBody)
if err != nil {
return fmt.Errorf("parse batch result: %w", err)
return recordPollFailureForTasks(ctx, adaptor, tasks, pollClassHookError, resp.StatusCode, err.Error())
}
for upstreamID, responseItem := range responseItems {
if ctx.Err() != nil {
......@@ -287,7 +312,27 @@ func updateBatchTasks(ctx context.Context, adaptor BatchTaskPollingAdaptor, chan
continue
}
snap := task.Snapshot()
task.Status = lo.If(model.TaskStatus(responseItem.TaskInfo.Status) != "", model.TaskStatus(responseItem.TaskInfo.Status)).Else(task.Status)
httpClass := classifyPollHTTP(resp.StatusCode)
parsedStatus := model.TaskStatus(responseItem.TaskInfo.Status)
if parsedStatus == model.TaskStatusUnknown || parsedStatus == "" || !knownPollStatus(parsedStatus) {
if err := recordPollFailure(ctx, adaptor, task, snap.Status, pollClassUnrecognized, resp.StatusCode, responseItem.TaskInfo.Reason); err != nil {
common.SysLog("UpdateSunoTask task error: " + err.Error())
}
continue
}
if httpClass == pollClassOtherClient && isNonTerminalPollStatus(parsedStatus) {
if err := recordPollFailure(ctx, adaptor, task, snap.Status, pollClassUnrecognized, resp.StatusCode, responseItem.TaskInfo.Reason); err != nil {
common.SysLog("UpdateSunoTask task error: " + err.Error())
}
continue
}
if isNonTerminalPollStatus(parsedStatus) {
task.PrivateData.PollFailures = 0
}
if len(responseItem.TaskInfo.PluginState) > 0 {
task.PrivateData.PluginState = responseItem.TaskInfo.PluginState
}
task.Status = lo.If(parsedStatus != "", parsedStatus).Else(task.Status)
task.FailReason = lo.If(responseItem.TaskInfo.Reason != "", responseItem.TaskInfo.Reason).Else(task.FailReason)
task.SubmitTime = lo.If(responseItem.SubmitTime != 0, responseItem.SubmitTime).Else(task.SubmitTime)
task.StartTime = lo.If(responseItem.StartTime != 0, responseItem.StartTime).Else(task.StartTime)
......@@ -447,24 +492,28 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch *
if privateData.Key != "" {
key = privateData.Key
}
resp, err := adaptor.FetchTask(baseURL, key, map[string]any{
"task_id": task.GetUpstreamTaskID(),
"action": constant.NormalizeTaskAction(task.Action),
"model": task.Properties.OriginModelName,
"upstream_model": task.Properties.UpstreamModelName,
}, proxy)
snap := task.Snapshot()
resp, err := adaptor.FetchTask(baseURL, key, task, proxy)
if err != nil {
return fmt.Errorf("fetchTask failed for task %s: %w", taskId, err)
return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassTransport, 0, err.Error())
}
defer resp.Body.Close()
responseBody, err := io.ReadAll(resp.Body)
if err != nil {
return fmt.Errorf("readAll failed for task %s: %w", taskId, err)
return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassTransport, resp.StatusCode, err.Error())
}
logger.LogDebug(ctx, "updateVideoSingleTask response: %s", responseBody)
snap := task.Snapshot()
switch classifyPollHTTP(resp.StatusCode) {
case pollClassNotFound:
return failTaskFromPoll(ctx, adaptor, task, snap.Status, fmt.Sprintf("upstream task not found (HTTP %d)", resp.StatusCode))
case pollClassAuth:
logger.LogWarn(ctx, fmt.Sprintf("task poll auth failure channel_id=%d task=%s http=%d", ch.Id, task.TaskID, resp.StatusCode))
return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassAuth, resp.StatusCode, "")
case pollClassTransient:
return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassTransient, resp.StatusCode, "")
}
taskResult := &relaycommon.TaskInfo{}
// try parse as New API response format
......@@ -478,41 +527,33 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch *
taskResult.Progress = t.Progress
taskResult.Reason = t.FailReason
task.Data = t.Data
} else if taskResult, err = adaptor.ParseTaskResult(responseBody); err != nil {
return fmt.Errorf("parseTaskResult failed for task %s: %w", taskId, err)
} else if taskResult, err = adaptor.ParseTaskResult(task, resp, responseBody); err != nil {
return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassHookError, resp.StatusCode, err.Error())
}
task.Data = redactVideoResponseBody(responseBody)
logger.LogDebug(ctx, "updateVideoSingleTask taskResult: %+v", taskResult)
now := time.Now().Unix()
if taskResult.Status == "" {
//taskResult = relaycommon.FailTaskInfo("upstream returned empty status")
errorResult := &dto.GeneralErrorResponse{}
if err = common.Unmarshal(responseBody, &errorResult); err == nil {
openaiError := errorResult.TryToOpenAIError()
if openaiError != nil {
// 返回规范的 OpenAI 错误格式,提取错误信息,判断错误是否为任务失败
if openaiError.Code == "429" {
// 429 错误通常表示请求过多或速率限制,暂时不认为是任务失败,保持原状态等待下一轮轮询
return nil
}
parsedStatus := model.TaskStatus(taskResult.Status)
if parsedStatus == model.TaskStatusUnknown || parsedStatus == "" || !knownPollStatus(parsedStatus) {
return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassUnrecognized, resp.StatusCode, unrecognizedPollDetail(taskResult.Reason, responseBody))
}
if classifyPollHTTP(resp.StatusCode) == pollClassOtherClient && isNonTerminalPollStatus(parsedStatus) {
return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassUnrecognized, resp.StatusCode, unrecognizedPollDetail(taskResult.Reason, responseBody))
}
// 其他错误认为是任务失败,记录错误信息并更新任务状态
taskResult = relaycommon.FailTaskInfo("upstream returned error")
} else {
// unknown error format, log original response
logger.LogError(ctx, fmt.Sprintf("Task %s returned empty status with unrecognized error format, response: %s", taskId, string(responseBody)))
taskResult = relaycommon.FailTaskInfo("upstream returned unrecognized message")
}
}
task.Data = redactVideoResponseBody(responseBody)
if len(taskResult.PluginState) > 0 {
task.PrivateData.PluginState = taskResult.PluginState
}
if isNonTerminalPollStatus(parsedStatus) {
task.PrivateData.PollFailures = 0
}
now := time.Now().Unix()
shouldFinalizeBilling := false
task.Status = model.TaskStatus(taskResult.Status)
switch taskResult.Status {
task.Status = parsedStatus
switch parsedStatus {
case model.TaskStatusSubmitted:
task.Progress = taskcommon.ProgressSubmitted
case model.TaskStatusQueued:
......@@ -549,8 +590,6 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch *
logger.LogInfo(ctx, fmt.Sprintf("Task %s failed: %s", task.TaskID, task.FailReason))
taskResult.Progress = taskcommon.ProgressComplete
shouldFinalizeBilling = true
default:
return fmt.Errorf("unknown task status %s for task %s", taskResult.Status, task.TaskID)
}
if taskResult.Progress != "" {
task.Progress = taskResult.Progress
......@@ -670,3 +709,131 @@ func settleTaskBillingOnComplete(ctx context.Context, adaptor TaskPollingAdaptor
}
return false
}
func classifyPollHTTP(statusCode int) string {
switch {
case statusCode >= 200 && statusCode < 300:
return pollClassOK
case statusCode == http.StatusNotFound || statusCode == http.StatusGone:
return pollClassNotFound
case statusCode == http.StatusUnauthorized || statusCode == http.StatusForbidden:
return pollClassAuth
case statusCode == http.StatusTooManyRequests || statusCode >= 500:
return pollClassTransient
case statusCode >= 400 && statusCode < 500:
return pollClassOtherClient
default:
return pollClassTransient
}
}
func knownPollStatus(status model.TaskStatus) bool {
switch status {
case model.TaskStatusNotStart, model.TaskStatusSubmitted, model.TaskStatusQueued, model.TaskStatusInProgress, model.TaskStatusSuccess, model.TaskStatusFailure:
return true
default:
return false
}
}
func isNonTerminalPollStatus(status model.TaskStatus) bool {
switch status {
case model.TaskStatusNotStart, model.TaskStatusSubmitted, model.TaskStatusQueued, model.TaskStatusInProgress:
return true
default:
return false
}
}
func pollFailureReason(class string, statusCode int, detail string) string {
reason := fmt.Sprintf("poll failed: %s", class)
if statusCode > 0 {
reason = fmt.Sprintf("poll failed: %s (HTTP %d)", class, statusCode)
}
if detail != "" {
reason = reason + ": " + detail
}
return reason
}
// unrecognizedPollDetail pairs the plugin's reason with a bounded copy of the
// upstream body so the WARN line is enough to diagnose a parser gap.
func unrecognizedPollDetail(reason string, body []byte) string {
const maxBodyChars = 512
redacted := string(redactVideoResponseBody(body))
if len(redacted) > maxBodyChars {
redacted = redacted[:maxBodyChars] + "…"
}
if strings.TrimSpace(reason) == "" {
return "body=" + redacted
}
return reason + "; body=" + redacted
}
func recordPollFailure(ctx context.Context, adaptor TaskPollingAdaptor, task *model.Task, fromStatus model.TaskStatus, class string, statusCode int, detail string) error {
task.PrivateData.PollFailures++
if class == pollClassUnrecognized || class == pollClassHookError {
// The redacted body is intentionally not persisted to Task.Data on these
// paths, so the WARN line is the only operator-visible copy of what the
// plugin could not interpret.
logger.LogWarn(ctx, fmt.Sprintf("task %s poll %s (failures=%d, http=%d): %s", task.TaskID, class, task.PrivateData.PollFailures, statusCode, detail))
}
// TASK_POLL_MAX_FAILURES <= 0 disables the consecutive-failure cutoff, matching
// TASK_TIMEOUT_MINUTES semantics; the 24h sweep remains the only backstop.
if constant.TaskPollMaxFailures > 0 && task.PrivateData.PollFailures >= constant.TaskPollMaxFailures {
return failTaskFromPoll(ctx, adaptor, task, fromStatus, pollFailureReason(class, statusCode, detail))
}
if _, err := task.UpdateWithStatus(fromStatus); err != nil {
return err
}
return nil
}
func recordPollFailureForTasks(ctx context.Context, adaptor TaskPollingAdaptor, tasks []*model.Task, class string, statusCode int, detail string) error {
var firstErr error
for _, task := range tasks {
if task == nil {
continue
}
if err := recordPollFailure(ctx, adaptor, task, task.Status, class, statusCode, detail); err != nil && firstErr == nil {
firstErr = err
}
}
return firstErr
}
func failTaskFromPoll(ctx context.Context, adaptor TaskPollingAdaptor, task *model.Task, fromStatus model.TaskStatus, reason string) error {
now := time.Now().Unix()
task.Status = model.TaskStatusFailure
task.Progress = taskcommon.ProgressComplete
if task.FinishTime == 0 {
task.FinishTime = now
}
task.FailReason = reason
won, err := task.UpdateWithStatus(fromStatus)
if err != nil {
return err
}
if !won {
return nil
}
taskResult := relaycommon.FailTaskInfo(reason)
billingSettled := settleTaskBillingOnComplete(ctx, adaptor, task, taskResult)
if !billingSettled && task.Quota != 0 {
RefundTaskQuota(ctx, task, reason)
}
return nil
}
func failTasksFromPoll(ctx context.Context, adaptor TaskPollingAdaptor, tasks []*model.Task, reason string) error {
var firstErr error
for _, task := range tasks {
if task == nil {
continue
}
if err := failTaskFromPoll(ctx, adaptor, task, task.Status, reason); err != nil && firstErr == nil {
firstErr = err
}
}
return firstErr
}
......@@ -40,12 +40,15 @@ type batchPollingAdaptor struct {
}
func (a *batchPollingAdaptor) FetchMode() string { return "batch" }
func (a *batchPollingAdaptor) FetchBatchTasks(_ string, _ string, taskIDs []string, _ string) (*http.Response, error) {
func (a *batchPollingAdaptor) FetchBatchTasks(_ string, _ string, tasks []*model.Task, _ string) (*http.Response, error) {
a.batchCalls++
a.batchIDs = append([]string(nil), taskIDs...)
a.batchIDs = a.batchIDs[:0]
for _, task := range tasks {
a.batchIDs = append(a.batchIDs, task.GetUpstreamTaskID())
}
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(bytes.NewReader([]byte(`{}`)))}, nil
}
func (a *batchPollingAdaptor) ParseBatchResult([]byte) (map[string]*BatchTaskResult, error) {
func (a *batchPollingAdaptor) ParseBatchResult(_ []*model.Task, _ *http.Response, _ []byte) (map[string]*BatchTaskResult, error) {
if a.results != nil {
return a.results, nil
}
......@@ -58,8 +61,11 @@ func (a *batchPollingAdaptor) ParseBatchResult([]byte) (map[string]*BatchTaskRes
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)
func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, task *model.Task, _ string) (*http.Response, error) {
taskID := ""
if task != nil {
taskID = task.GetUpstreamTaskID()
}
if taskID == a.blockTaskID && a.releaseBlock != nil {
a.blockOnce.Do(func() {
if a.blockStarted != nil {
......@@ -97,7 +103,7 @@ func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, body map[string]
}, nil
}
func (a *taskPollingFetchAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) {
func (a *taskPollingFetchAdaptor) ParseTaskResult(*model.Task, *http.Response, []byte) (*relaycommon.TaskInfo, error) {
return &relaycommon.TaskInfo{Status: model.TaskStatusInProgress}, nil
}
......@@ -721,3 +727,273 @@ func TestSweepTimedOutTasksHonorsRefundRolloutBoundary(t *testing.T) {
assert.Equal(t, initialQuota+modernTaskQuota, getUserQuota(t, userID))
assert.Equal(t, int64(1), countLogs(t))
}
type scriptedPollingAdaptor struct {
statusCode int
body []byte
fetchErr error
parse *relaycommon.TaskInfo
parseErr error
}
func (a *scriptedPollingAdaptor) Init(*relaycommon.RelayInfo) {}
func (a *scriptedPollingAdaptor) FetchTask(string, string, *model.Task, string) (*http.Response, error) {
if a.fetchErr != nil {
return nil, a.fetchErr
}
code := a.statusCode
if code == 0 {
code = http.StatusOK
}
body := a.body
if body == nil {
body = []byte(`{}`)
}
return &http.Response{StatusCode: code, Body: io.NopCloser(bytes.NewReader(body))}, nil
}
func (a *scriptedPollingAdaptor) ParseTaskResult(*model.Task, *http.Response, []byte) (*relaycommon.TaskInfo, error) {
if a.parseErr != nil {
return nil, a.parseErr
}
if a.parse != nil {
return a.parse, nil
}
return &relaycommon.TaskInfo{Status: model.TaskStatusInProgress}, nil
}
func (a *scriptedPollingAdaptor) AdjustBillingOnComplete(*model.Task, *relaycommon.TaskInfo) int {
return 0
}
type scriptedBatchPollingAdaptor struct {
scriptedPollingAdaptor
results map[string]*BatchTaskResult
}
func (a *scriptedBatchPollingAdaptor) FetchMode() string { return "batch" }
func (a *scriptedBatchPollingAdaptor) FetchBatchTasks(string, string, []*model.Task, string) (*http.Response, error) {
return a.FetchTask("", "", nil, "")
}
func (a *scriptedBatchPollingAdaptor) ParseBatchResult([]*model.Task, *http.Response, []byte) (map[string]*BatchTaskResult, error) {
if a.parseErr != nil {
return nil, a.parseErr
}
return a.results, nil
}
func TestUpdateVideoSingleTaskPollClassification(t *testing.T) {
testCases := []struct {
name string
statusCode int
fetchErr error
parse *relaycommon.TaskInfo
parseErr error
priorFailures int
priorState string
maxFailures int
wantStatus model.TaskStatus
wantFailures int
wantRefund bool
wantReason string
wantState string
wantUnchanged bool
}{
{
name: "404 fails immediately and refunds",
statusCode: http.StatusNotFound,
wantStatus: model.TaskStatusFailure,
wantRefund: true,
wantReason: "upstream task not found (HTTP 404)",
wantUnchanged: false,
},
{
name: "401 increments without changing status",
statusCode: http.StatusUnauthorized,
wantStatus: model.TaskStatusInProgress,
wantFailures: 1,
wantUnchanged: true,
},
{
name: "429 reaches threshold and refunds",
statusCode: http.StatusTooManyRequests,
priorFailures: 2,
maxFailures: 3,
wantStatus: model.TaskStatusFailure,
wantFailures: 3,
wantRefund: true,
wantReason: "poll failed: transient (HTTP 429)",
},
{
name: "UNKNOWN increments",
statusCode: http.StatusOK,
parse: &relaycommon.TaskInfo{Status: model.TaskStatusUnknown, Reason: "weird"},
wantStatus: model.TaskStatusInProgress,
wantFailures: 1,
wantUnchanged: true,
},
{
name: "valid 2xx resets the failure counter",
statusCode: http.StatusOK,
parse: &relaycommon.TaskInfo{Status: model.TaskStatusInProgress},
priorFailures: 5,
wantStatus: model.TaskStatusInProgress,
wantFailures: 0,
},
{
name: "omit state preserves previous plugin state",
statusCode: http.StatusOK,
parse: &relaycommon.TaskInfo{Status: model.TaskStatusInProgress},
priorState: `{"req_key":"keep"}`,
wantStatus: model.TaskStatusInProgress,
wantState: `{"req_key":"keep"}`,
},
{
name: "returned state replaces plugin state",
statusCode: http.StatusOK,
parse: &relaycommon.TaskInfo{
Status: model.TaskStatusInProgress,
PluginState: []byte(`{"req_key":"new"}`),
},
priorState: `{"req_key":"old"}`,
wantStatus: model.TaskStatusInProgress,
wantState: `{"req_key":"new"}`,
},
{
name: "other 4xx non-terminal is unrecognized",
statusCode: http.StatusBadRequest,
parse: &relaycommon.TaskInfo{Status: model.TaskStatusInProgress},
wantStatus: model.TaskStatusInProgress,
wantFailures: 1,
wantUnchanged: true,
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
truncate(t)
const userID, tokenID, channelID = 510, 510, 510
const initialQuota, preConsumed, tokenRemain = 10_000, 4_000, 7_000
seedUser(t, userID, initialQuota)
seedToken(t, tokenID, userID, "sk-poll-class", tokenRemain)
ch := &model.Channel{Id: channelID, Type: constant.ChannelTypeKling, Name: "poll", Key: "sk-test", Status: common.ChannelStatusEnabled}
if testCase.maxFailures > 0 {
previous := constant.TaskPollMaxFailures
constant.TaskPollMaxFailures = testCase.maxFailures
t.Cleanup(func() { constant.TaskPollMaxFailures = previous })
}
task := makeTask(userID, channelID, preConsumed, tokenID, BillingSourceWallet, 0)
task.TaskID = "task_poll_class"
task.PrivateData.UpstreamTaskID = "upstream_poll_class"
task.PrivateData.PollFailures = testCase.priorFailures
if testCase.priorState != "" {
task.PrivateData.PluginState = []byte(testCase.priorState)
}
require.NoError(t, model.DB.Create(task).Error)
adaptor := &scriptedPollingAdaptor{statusCode: testCase.statusCode, fetchErr: testCase.fetchErr, parse: testCase.parse, parseErr: testCase.parseErr}
require.NoError(t, updateVideoSingleTask(context.Background(), adaptor, ch, task.GetUpstreamTaskID(), map[string]*model.Task{
task.GetUpstreamTaskID(): task,
}))
var persisted model.Task
require.NoError(t, model.DB.First(&persisted, task.ID).Error)
assert.EqualValues(t, testCase.wantStatus, persisted.Status)
assert.Equal(t, testCase.wantFailures, persisted.PrivateData.PollFailures)
if testCase.wantUnchanged {
assert.Empty(t, persisted.FailReason)
}
if testCase.wantReason != "" {
assert.Contains(t, persisted.FailReason, testCase.wantReason)
}
if testCase.wantState != "" {
assert.JSONEq(t, testCase.wantState, string(persisted.PrivateData.PluginState))
}
if testCase.wantRefund {
assert.Equal(t, initialQuota+preConsumed, getUserQuota(t, userID))
assert.Equal(t, tokenRemain+preConsumed, getTokenRemainQuota(t, tokenID))
assert.Zero(t, persisted.Quota)
log := getLastLog(t)
require.NotNil(t, log)
assert.Equal(t, model.LogTypeRefund, log.Type)
} else {
assert.Equal(t, initialQuota, getUserQuota(t, userID))
assert.Equal(t, tokenRemain, getTokenRemainQuota(t, tokenID))
}
})
}
}
func TestUpdateBatchTasksPollClassification(t *testing.T) {
testCases := []struct {
name string
statusCode int
resultStatus model.TaskStatus
wantStatus model.TaskStatus
wantFailures int
wantRefund bool
wantReason string
}{
{
name: "404 fails the batch and refunds",
statusCode: http.StatusNotFound,
wantStatus: model.TaskStatusFailure,
wantRefund: true,
wantReason: "upstream task not found (HTTP 404)",
},
{
name: "401 increments every task",
statusCode: http.StatusUnauthorized,
wantStatus: model.TaskStatusInProgress,
wantFailures: 1,
},
{
name: "UNKNOWN increments",
statusCode: http.StatusOK,
resultStatus: model.TaskStatusUnknown,
wantStatus: model.TaskStatusInProgress,
wantFailures: 1,
},
}
for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
truncate(t)
const userID, tokenID, channelID = 610, 610, 610
const initialQuota, preConsumed, tokenRemain = 10_000, 4_000, 7_000
seedUser(t, userID, initialQuota)
seedToken(t, tokenID, userID, "sk-batch-class", tokenRemain)
seedTaskPollingChannel(t, channelID, true)
task := makeTask(userID, channelID, preConsumed, tokenID, BillingSourceWallet, 0)
task.TaskID = "task_batch_class"
task.PrivateData.UpstreamTaskID = "upstream_batch_class"
require.NoError(t, model.DB.Create(task).Error)
upstreamID := task.GetUpstreamTaskID()
adaptor := &scriptedBatchPollingAdaptor{
scriptedPollingAdaptor: scriptedPollingAdaptor{statusCode: testCase.statusCode},
}
if testCase.resultStatus != "" {
adaptor.results = map[string]*BatchTaskResult{
upstreamID: {TaskInfo: relaycommon.TaskInfo{TaskID: upstreamID, Status: string(testCase.resultStatus), Reason: "weird"}},
}
}
require.NoError(t, UpdateBatchTasks(context.Background(), adaptor, map[int][]string{channelID: {upstreamID}}, map[string]*model.Task{upstreamID: task}))
var persisted model.Task
require.NoError(t, model.DB.First(&persisted, task.ID).Error)
assert.EqualValues(t, testCase.wantStatus, persisted.Status)
assert.Equal(t, testCase.wantFailures, persisted.PrivateData.PollFailures)
if testCase.wantReason != "" {
assert.Contains(t, persisted.FailReason, testCase.wantReason)
}
if testCase.wantRefund {
assert.Equal(t, initialQuota+preConsumed, getUserQuota(t, userID))
assert.Zero(t, persisted.Quota)
} else {
assert.Equal(t, initialQuota, getUserQuota(t, userID))
}
})
}
}
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or sign in to comment