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 @@ ...@@ -56,6 +56,10 @@
# 任务和功能配置 # 任务和功能配置
# 更新任务启用 # 更新任务启用
# UPDATE_TASK=true # UPDATE_TASK=true
# 异步任务硬超时(分钟),按提交时间计算,超时未完成的任务标记失败并退款;0 表示禁用
# TASK_TIMEOUT_MINUTES=1440
# 异步任务连续轮询失败阈值(上游 429/5xx/401/403、网络错误、无法识别的响应),达到后任务标记失败并退款;正常轮询成功一次即归零
# TASK_POLL_MAX_FAILURES=20
# 对话超时设置 # 对话超时设置
# 所有请求超时时间,单位秒,默认为0,表示不限制 # 所有请求超时时间,单位秒,默认为0,表示不限制
......
...@@ -202,6 +202,8 @@ func initConstantEnv() { ...@@ -202,6 +202,8 @@ func initConstantEnv() {
constant.TaskQueryLimit = GetEnvOrDefault("TASK_QUERY_LIMIT", 1000) constant.TaskQueryLimit = GetEnvOrDefault("TASK_QUERY_LIMIT", 1000)
// 异步任务超时时间(分钟),超过此时间未完成的任务将被标记为失败并退款。0 表示禁用。 // 异步任务超时时间(分钟),超过此时间未完成的任务将被标记为失败并退款。0 表示禁用。
constant.TaskTimeoutMinutes = GetEnvOrDefault("TASK_TIMEOUT_MINUTES", 1440) 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) constant.TaskPluginProtocolTimeoutSeconds = GetEnvOrDefault("TASK_PLUGIN_PROTOCOL_TIMEOUT_SECONDS", 600)
......
...@@ -18,6 +18,7 @@ var GenerateDefaultToken bool ...@@ -18,6 +18,7 @@ var GenerateDefaultToken bool
var ErrorLogEnabled bool var ErrorLogEnabled bool
var TaskQueryLimit int var TaskQueryLimit int
var TaskTimeoutMinutes int var TaskTimeoutMinutes int
var TaskPollMaxFailures = 20
var TaskPluginProtocolTimeoutSeconds int var TaskPluginProtocolTimeoutSeconds int
var TaskPluginProtocolTickMilliseconds int var TaskPluginProtocolTickMilliseconds int
var TaskPluginProtocolTickJitterMilliseconds int var TaskPluginProtocolTickJitterMilliseconds int
......
...@@ -366,14 +366,14 @@ type terminalSettlementPollingAdaptor struct { ...@@ -366,14 +366,14 @@ type terminalSettlementPollingAdaptor struct {
func (a *terminalSettlementPollingAdaptor) Init(*relaycommon.RelayInfo) {} 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{ return &http.Response{
StatusCode: http.StatusOK, StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(`{}`)), Body: io.NopCloser(strings.NewReader(`{}`)),
}, nil }, nil
} }
func (a *terminalSettlementPollingAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) { func (a *terminalSettlementPollingAdaptor) ParseTaskResult(*model.Task, *http.Response, []byte) (*relaycommon.TaskInfo, error) {
return &relaycommon.TaskInfo{ return &relaycommon.TaskInfo{
Status: model.TaskStatusSuccess, Status: model.TaskStatusSuccess,
Progress: "100%", Progress: "100%",
......
...@@ -744,6 +744,9 @@ func executeTaskSubmissionWith( ...@@ -744,6 +744,9 @@ func executeTaskSubmissionWith(
} }
task.Quota = result.Quota task.Quota = result.Quota
task.Data = result.TaskData task.Data = result.TaskData
if len(result.PluginState) > 0 {
task.PrivateData.PluginState = result.PluginState
}
task.Action = relayInfo.Action task.Action = relayInfo.Action
if immediate := result.Immediate; immediate != nil { if immediate := result.Immediate; immediate != nil {
task.Status = model.TaskStatus(immediate.Status) task.Status = model.TaskStatus(immediate.Status)
......
...@@ -5,6 +5,7 @@ import ( ...@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"regexp"
"strings" "strings"
"testing" "testing"
...@@ -117,6 +118,17 @@ func TestDisableThirdPartyPluginSupportsCascadeAndForce(t *testing.T) { ...@@ -117,6 +118,17 @@ func TestDisableThirdPartyPluginSupportsCascadeAndForce(t *testing.T) {
assert.Equal(t, common.ChannelStatusManuallyDisabled, updated.Status) 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) { func setupTaskPluginFactoryDisableTest(t *testing.T) {
t.Helper() t.Helper()
setupTaskPluginControllerTest(t) setupTaskPluginControllerTest(t)
...@@ -240,7 +252,9 @@ func TestDisableFactoryOverrideRowKeepsEnabledFlagPath(t *testing.T) { ...@@ -240,7 +252,9 @@ func TestDisableFactoryOverrideRowKeepsEnabledFlagPath(t *testing.T) {
setupTaskPluginFactoryDisableTest(t) setupTaskPluginFactoryDisableTest(t)
factorySource, err := plugins.Source("kling") factorySource, err := plugins.Source("kling")
require.NoError(t, err) 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{}) loaded, err := jsplugin.DefaultRegistry.Register(overrideSource, jsplugin.Options{})
require.NoError(t, err) require.NoError(t, err)
t.Cleanup(func() { jsplugin.DefaultRegistry.Unregister("kling") }) t.Cleanup(func() { jsplugin.DefaultRegistry.Unregister("kling") })
...@@ -266,7 +280,7 @@ func TestDisableFactoryOverrideRowKeepsEnabledFlagPath(t *testing.T) { ...@@ -266,7 +280,7 @@ func TestDisableFactoryOverrideRowKeepsEnabledFlagPath(t *testing.T) {
assert.True(t, taskPluginOptionsHasKey(t, "kling")) assert.True(t, taskPluginOptionsHasKey(t, "kling"))
got, ok := jsplugin.DefaultRegistry.Get("kling") got, ok := jsplugin.DefaultRegistry.Get("kling")
require.True(t, ok) 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) { func TestListTaskPluginsIncludesFactoryWithoutDatabaseRows(t *testing.T) {
...@@ -358,7 +372,9 @@ func TestListTaskPluginsShowsDisabledFallbackWhenOverridesAreDisabled(t *testing ...@@ -358,7 +372,9 @@ func TestListTaskPluginsShowsDisabledFallbackWhenOverridesAreDisabled(t *testing
setupTaskPluginControllerTest(t) setupTaskPluginControllerTest(t)
factorySource, err := plugins.Source("kling") factorySource, err := plugins.Source("kling")
require.NoError(t, err) 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{}) loaded, err := jsplugin.DefaultRegistry.Register(overrideSource, jsplugin.Options{})
require.NoError(t, err) require.NoError(t, err)
plugin := model.TaskPlugin{ plugin := model.TaskPlugin{
...@@ -399,7 +415,9 @@ func TestDeleteActiveOverrideFallsBackToFactoryAndDeletesRecord(t *testing.T) { ...@@ -399,7 +415,9 @@ func TestDeleteActiveOverrideFallsBackToFactoryAndDeletesRecord(t *testing.T) {
setupTaskPluginControllerTest(t) setupTaskPluginControllerTest(t)
factorySource, err := plugins.Source("kling") factorySource, err := plugins.Source("kling")
require.NoError(t, err) 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"}) loaded, err := jsplugin.DefaultRegistry.Register(overrideSource, jsplugin.Options{Key: "kling", Version: "test-override"})
require.NoError(t, err) require.NoError(t, err)
t.Cleanup(func() { jsplugin.DefaultRegistry.Unregister("kling") }) t.Cleanup(func() { jsplugin.DefaultRegistry.Unregister("kling") })
......
...@@ -25,6 +25,9 @@ export type UsageExample = {label: string; facts: Readonly<Record<string, string ...@@ -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 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 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 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 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 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} 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: { ...@@ -36,13 +39,13 @@ export declare const protocols: {
openai_video?: {decodeRequest(ctx: ProtocolDecodeContext): SubmitIntent; render(ctx: unknown, task: TaskView): unknown}; openai_video?: {decodeRequest(ctx: ProtocolDecodeContext): SubmitIntent; render(ctx: unknown, task: TaskView): unknown};
}; };
export declare function buildSubmitRequest(ctx: DriverContext): RequestDescriptor; export declare function buildSubmitRequest(ctx: DriverContext): RequestDescriptor;
export declare function parseSubmitResponse(ctx: DriverContext, response: UpstreamResponse): {taskId: string; taskData?: unknown; immediate?: NormalizedTaskResult}; export declare function parseSubmitResponse(ctx: DriverContext, response: UpstreamResponse): {taskId: string; taskData?: unknown; immediate?: NormalizedTaskResult; state?: unknown};
export declare function buildQueryRequest(ctx: DriverContext & {taskId: string}): RequestDescriptor; export declare function buildQueryRequest(ctx: TaskQueryContext): RequestDescriptor;
export declare function buildBatchQueryRequest(ctx: DriverContext, taskIds: readonly string[]): RequestDescriptor; export declare function buildBatchQueryRequest(ctx: BatchQueryContext, tasks: readonly TaskQueryContext[]): RequestDescriptor;
export declare function parseTaskResult(ctx: DriverContext, body: unknown): NormalizedTaskResult; export declare function parseTaskResult(ctx: TaskQueryContext, body: unknown, response: HookHTTPResponse): NormalizedTaskResult;
export declare function parseBatchResult(ctx: DriverContext, body: unknown): readonly (NormalizedTaskResult & {taskId: string; data?: unknown})[]; 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 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 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 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 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 ...@@ -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. 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. `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 ` ...@@ -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. 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. `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 @@ ...@@ -41,6 +41,45 @@
"maxBytes": {"type": "integer", "exclusiveMinimum": 0} "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 { ...@@ -140,9 +140,7 @@ func (o *LogOther) toMap() map[string]any {
return result return result
} }
for key, value := range o.public { maps.Copy(result, o.public)
result[key] = value
}
if adminInfo := copyLogOtherMap(o.adminInfo); len(adminInfo) > 0 { if adminInfo := copyLogOtherMap(o.adminInfo); len(adminInfo) > 0 {
result[logOtherAdminInfoKey] = adminInfo result[logOtherAdminInfoKey] = adminInfo
} }
......
...@@ -126,6 +126,11 @@ type TaskPrivateData struct { ...@@ -126,6 +126,11 @@ type TaskPrivateData struct {
// disconnect regardless; this only echoes the protocol-level request // disconnect regardless; this only echoes the protocol-level request
// attribute back on retrieval snapshots. // attribute back on retrieval snapshots.
ResponsesBackground bool `json:"responses_background,omitempty"` 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 { type TaskExecutionSnapshot struct {
...@@ -194,7 +199,10 @@ func (p *TaskPrivateData) Scan(val interface{}) error { ...@@ -194,7 +199,10 @@ func (p *TaskPrivateData) Scan(val interface{}) error {
} }
func (p TaskPrivateData) Value() (driver.Value, 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 return nil, nil
} }
// 同 Properties.Value:string 避免 PG simple protocol 的 bytea 编码。 // 同 Properties.Value:string 避免 PG simple protocol 的 bytea 编码。
...@@ -466,13 +474,15 @@ func (Task *Task) InsertWithContext(ctx context.Context) error { ...@@ -466,13 +474,15 @@ func (Task *Task) InsertWithContext(ctx context.Context) error {
} }
type taskSnapshot struct { type taskSnapshot struct {
Status TaskStatus Status TaskStatus
Progress string Progress string
StartTime int64 StartTime int64
FinishTime int64 FinishTime int64
FailReason string FailReason string
ResultURL string ResultURL string
Data json.RawMessage Data json.RawMessage
PluginState json.RawMessage
PollFailures int
} }
func (s taskSnapshot) Equal(other taskSnapshot) bool { func (s taskSnapshot) Equal(other taskSnapshot) bool {
...@@ -482,18 +492,22 @@ func (s taskSnapshot) Equal(other taskSnapshot) bool { ...@@ -482,18 +492,22 @@ func (s taskSnapshot) Equal(other taskSnapshot) bool {
s.FinishTime == other.FinishTime && s.FinishTime == other.FinishTime &&
s.FailReason == other.FailReason && s.FailReason == other.FailReason &&
s.ResultURL == other.ResultURL && 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 { func (t *Task) Snapshot() taskSnapshot {
return taskSnapshot{ return taskSnapshot{
Status: t.Status, Status: t.Status,
Progress: t.Progress, Progress: t.Progress,
StartTime: t.StartTime, StartTime: t.StartTime,
FinishTime: t.FinishTime, FinishTime: t.FinishTime,
FailReason: t.FailReason, FailReason: t.FailReason,
ResultURL: t.PrivateData.ResultURL, ResultURL: t.PrivateData.ResultURL,
Data: t.Data, Data: t.Data,
PluginState: t.PrivateData.PluginState,
PollFailures: t.PrivateData.PollFailures,
} }
} }
......
...@@ -177,6 +177,29 @@ func TestSnapshotEqual_NilVsEmpty(t *testing.T) { ...@@ -177,6 +177,29 @@ func TestSnapshotEqual_NilVsEmpty(t *testing.T) {
assert.True(t, a.Equal(b)) 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) { func TestSnapshot_Roundtrip(t *testing.T) {
task := &Task{ task := &Task{
Status: TaskStatusInProgress, Status: TaskStatusInProgress,
...@@ -185,7 +208,9 @@ func TestSnapshot_Roundtrip(t *testing.T) { ...@@ -185,7 +208,9 @@ func TestSnapshot_Roundtrip(t *testing.T) {
FinishTime: 5678, FinishTime: 5678,
FailReason: "timeout", FailReason: "timeout",
PrivateData: TaskPrivateData{ 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"}`), Data: json.RawMessage(`{"model":"test-model"}`),
} }
...@@ -197,6 +222,8 @@ func TestSnapshot_Roundtrip(t *testing.T) { ...@@ -197,6 +222,8 @@ func TestSnapshot_Roundtrip(t *testing.T) {
assert.Equal(t, task.FailReason, snap.FailReason) assert.Equal(t, task.FailReason, snap.FailReason)
assert.Equal(t, task.PrivateData.ResultURL, snap.ResultURL) assert.Equal(t, task.PrivateData.ResultURL, snap.ResultURL)
assert.JSONEq(t, string(task.Data), string(snap.Data)) 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) { ...@@ -292,3 +319,30 @@ func TestUpdateWithStatus_ConcurrentWinner(t *testing.T) {
} }
assert.Equal(t, 1, winCount, "exactly one goroutine should win the CAS") 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 { ...@@ -33,6 +33,7 @@ func RunCLI(args []string, stdout, stderr io.Writer) int {
fmt.Fprintf(stderr, "plugin lint failed: %v\n", compileErr) fmt.Fprintf(stderr, "plugin lint failed: %v\n", compileErr)
return 1 return 1
} }
warnParseTaskResultInProgressFallback(string(source), stderr)
fmt.Fprintf(stdout, "plugin %s@%s is valid\n", plugin.Meta.Key, plugin.Meta.Version) fmt.Fprintf(stdout, "plugin %s@%s is valid\n", plugin.Meta.Key, plugin.Meta.Version)
return 0 return 0
case "test": case "test":
...@@ -57,3 +58,35 @@ func RunCLI(args []string, stdout, stderr io.Writer) int { ...@@ -57,3 +58,35 @@ func RunCLI(args []string, stdout, stderr io.Writer) int {
return 2 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) { ...@@ -28,6 +28,25 @@ func TestPluginCLI(t *testing.T) {
assert.Contains(t, stdout.String(), "1/1 cases") 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 = ` 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 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}; } export function buildSubmitRequest(ctx) { return {url: ctx.baseUrl}; }
......
...@@ -346,6 +346,9 @@ func TestHailuoParseTaskResult(t *testing.T) { ...@@ -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"}, {"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 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", "", ""}, {"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 { for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { t.Run(testCase.name, func(t *testing.T) {
...@@ -451,7 +454,7 @@ func TestHailuoH3CompletionUsageFacts(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) { t.Run("polling adaptor carries actual facts into task settlement", func(t *testing.T) {
adaptor := taskplugin.New(plugin) 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}}}`, `{"task":{"id":"1","status":"succeeded","resolution":"2K","usage":{"output_seconds":5,"input_seconds":7.5,"input_image_count":6}}}`,
)) ))
require.NoError(t, err) require.NoError(t, err)
......
package plugins_test 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) { func TestJimengResponsesProtocol(t *testing.T) {
testVideoResponsesProtocol(t, videoResponsesTestCase{ testVideoResponsesProtocol(t, videoResponsesTestCase{
...@@ -25,3 +33,60 @@ func TestJimengResponsesProtocol(t *testing.T) { ...@@ -25,3 +33,60 @@ func TestJimengResponsesProtocol(t *testing.T) {
wantVendorName: "jimeng", 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 = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "Alibaba Cloud Bailian Wanxiang video generation (text-to-video and image-to-video)", en: "Alibaba Cloud Bailian Wanxiang video generation (text-to-video and image-to-video)",
zh: "阿里云百炼万相视频生成(文生视频、图生视频)", zh: "阿里云百炼万相视频生成(文生视频、图生视频)",
}, },
version: "1.0.0", version: "1.0.1",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [17], channelTypes: [17],
models: [ models: [
...@@ -270,7 +270,7 @@ export function parseTaskResult(ctx, body) { ...@@ -270,7 +270,7 @@ export function parseTaskResult(ctx, body) {
if (!reason) reason = "task failed"; if (!reason) reason = "task failed";
return { status: "FAILURE", reason: reason }; return { status: "FAILURE", reason: reason };
} }
return { status: "QUEUED" }; return { status: "UNKNOWN", reason: "unrecognized status: " + String(output.task_status || "") };
} }
function artifactData(ctx) { function artifactData(ctx) {
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "Volcengine Doubao Seedance video generation (text-to-video, image-to-video, and video-to-video)", en: "Volcengine Doubao Seedance video generation (text-to-video, image-to-video, and video-to-video)",
zh: "火山引擎豆包 Seedance 视频生成(文生视频、图生视频、视频生视频)", zh: "火山引擎豆包 Seedance 视频生成(文生视频、图生视频、视频生视频)",
}, },
version: "1.0.0", version: "1.0.1",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [54, 45], // VolcEngine-type channels serve Ark video models with the same wire format channelTypes: [54, 45], // VolcEngine-type channels serve Ark video models with the same wire format
models: [ models: [
...@@ -324,7 +324,7 @@ export function parseTaskResult(ctx, body) { ...@@ -324,7 +324,7 @@ export function parseTaskResult(ctx, body) {
const reason = body.error && body.error.message ? body.error.message : body.status; const reason = body.error && body.error.message ? body.error.message : body.status;
return { status: "FAILURE", progress: "100%", reason: reason }; 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) { function artifactData(ctx) {
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "Google Veo video generation on the Gemini API (text-to-video and image-to-video)", en: "Google Veo video generation on the Gemini API (text-to-video and image-to-video)",
zh: "Google Veo 视频生成(文生视频、图生视频),Gemini API 版本", zh: "Google Veo 视频生成(文生视频、图生视频),Gemini API 版本",
}, },
version: "1.0.0", version: "1.0.1",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [24], 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"], 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) { ...@@ -209,7 +209,11 @@ export function buildQueryRequest(ctx) {
export function parseTaskResult(ctx, body) { export function parseTaskResult(ctx, body) {
if (body.error && body.error.message) return { status: "FAILURE", progress: "100%", reason: body.error.message }; 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 videos = ((body.response || {}).generateVideoResponse || {}).generatedVideos || [];
const uri = videos.length && videos[0].video ? videos[0].video.uri || "" : ""; const uri = videos.length && videos[0].video ? videos[0].video.uri || "" : "";
return { taskId: utils.base64URL(body.name || ""), status: "SUCCESS", progress: "100%", remoteUrl: uri }; return { taskId: utils.base64URL(body.name || ""), status: "SUCCESS", progress: "100%", remoteUrl: uri };
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "MiniMax Hailuo video generation (text-to-video, image-to-video, and MiniMax-H3 multimodal reference)", en: "MiniMax Hailuo video generation (text-to-video, image-to-video, and MiniMax-H3 multimodal reference)",
zh: "MiniMax 海螺视频生成(文生视频、图生视频、MiniMax-H3 多模态参考生视频)", zh: "MiniMax 海螺视频生成(文生视频、图生视频、MiniMax-H3 多模态参考生视频)",
}, },
version: "1.1.1", version: "1.1.2",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [35], channelTypes: [35],
models: [ models: [
...@@ -465,8 +465,6 @@ export function buildQueryRequest(ctx) { ...@@ -465,8 +465,6 @@ export function buildQueryRequest(ctx) {
} }
export function parseTaskResult(ctx, body) { 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); const apiError = h3APIError(body);
if (apiError) { if (apiError) {
if (apiError.statusCode === 408 || apiError.statusCode === 429 || apiError.statusCode >= 500) throw new Error(apiError.message); if (apiError.statusCode === 408 || apiError.statusCode === 429 || apiError.statusCode >= 500) throw new Error(apiError.message);
...@@ -475,7 +473,10 @@ export function parseTaskResult(ctx, body) { ...@@ -475,7 +473,10 @@ export function parseTaskResult(ctx, body) {
const h3Task = h3QueryTask(body); const h3Task = h3QueryTask(body);
if (h3Task) { if (h3Task) {
const h3Statuses = { queued: "QUEUED", running: "IN_PROGRESS", succeeded: "SUCCESS", failed: "FAILURE", cancelled: "FAILURE" }; 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%" }; const h3Result = { code: 0, status: h3Status, progress: h3Status === "QUEUED" ? "30%" : h3Status === "IN_PROGRESS" ? "50%" : "100%" };
if (h3Status === "SUCCESS") { if (h3Status === "SUCCESS") {
const url = trimmed(h3Task.content && h3Task.content.url); const url = trimmed(h3Task.content && h3Task.content.url);
...@@ -486,11 +487,17 @@ export function parseTaskResult(ctx, body) { ...@@ -486,11 +487,17 @@ export function parseTaskResult(ctx, body) {
} }
return h3Result; 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 base = body.base_resp || {};
const statuses = { Preparing: "IN_PROGRESS", Queueing: "IN_PROGRESS", Processing: "IN_PROGRESS", Success: "SUCCESS", Fail: "FAILURE" }; 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 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 }; return { code: base.status_code || 0, status: status, progress: progress, reason: reason };
} }
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "Volcengine Jimeng video generation (text-to-video, image-to-video, and first-and-last-frame)", en: "Volcengine Jimeng video generation (text-to-video, image-to-video, and first-and-last-frame)",
zh: "火山引擎即梦视频生成(文生视频、图生视频、首尾帧)", zh: "火山引擎即梦视频生成(文生视频、图生视频、首尾帧)",
}, },
version: "1.0.0", version: "1.0.1",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [51], channelTypes: [51],
models: ["jimeng_vgfm_t2v_l20"], models: ["jimeng_vgfm_t2v_l20"],
...@@ -270,10 +270,8 @@ function filePlaceholder(image) { ...@@ -270,10 +270,8 @@ function filePlaceholder(image) {
} }
function queryReqKey(ctx) { function queryReqKey(ctx) {
const data = (ctx && ctx.data) || {}; const state = (ctx && ctx.state) || {};
if (typeof data.req_key === "string" && data.req_key.trim()) return data.req_key.trim(); if (typeof state.req_key === "string" && state.req_key.trim()) return state.req_key.trim();
const req = (ctx && ctx.requestBody) || {};
if (typeof req.req_key === "string" && req.req_key.trim()) return req.req_key.trim();
if (ctx && ctx.action === "image_to_video") return "jimeng_vgfm_i2v_l20"; 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"; if (ctx && ctx.action === "first_tail_to_video") return "jimeng_i2v_first_tail_v30";
return "jimeng_vgfm_t2v_l20"; return "jimeng_vgfm_t2v_l20";
...@@ -349,7 +347,7 @@ export function parseSubmitResponse(ctx, resp) { ...@@ -349,7 +347,7 @@ export function parseSubmitResponse(ctx, resp) {
const body = resp.body || {}; const body = resp.body || {};
if (body.code !== 10000) throw new Error(body.message || "jimeng submit failed"); 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"); 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) { export function extractUsage(ctx) {
...@@ -366,22 +364,20 @@ export function buildQueryRequest(ctx) { ...@@ -366,22 +364,20 @@ export function buildQueryRequest(ctx) {
export function parseTaskResult(ctx, body) { export function parseTaskResult(ctx, body) {
const data = body.data || {}; const data = body.data || {};
let status = "";
let progress = "";
if (body.code !== 10000) { if (body.code !== 10000) {
status = "FAILURE"; return { code: body.code || 0, status: "FAILURE", progress: "100%", reason: body.message || "" };
progress = "100%";
} }
if (data.status === "in_queue") { if (data.status === "in_queue") {
status = "QUEUED"; const result = { code: 0, status: "QUEUED", progress: "10%", reason: "" };
progress = "10%"; if (data.video_url) result.url = data.video_url;
} else if (data.status === "done") { return result;
status = "SUCCESS"; }
progress = "100%"; 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 || "" }; return { code: 0, status: "UNKNOWN", reason: "unrecognized status: " + String(data.status || "") };
if (data.video_url) result.url = data.video_url;
return result;
} }
function artifactData(ctx) { function artifactData(ctx) {
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "Kuaishou Kling video generation (text-to-video and image-to-video)", en: "Kuaishou Kling video generation (text-to-video and image-to-video)",
zh: "快手可灵视频生成(文生视频、图生视频)", zh: "快手可灵视频生成(文生视频、图生视频)",
}, },
version: "1.0.0", version: "1.0.1",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [50], channelTypes: [50],
models: ["kling-v1", "kling-v1-6", "kling-v2-master"], models: ["kling-v1", "kling-v1-6", "kling-v2-master"],
...@@ -286,7 +286,7 @@ export function parseTaskResult(ctx, body) { ...@@ -286,7 +286,7 @@ export function parseTaskResult(ctx, body) {
const data = body.data || {}; const data = body.data || {};
const statuses = { submitted: "SUBMITTED", processing: "IN_PROGRESS", succeed: "SUCCESS", failed: "FAILURE" }; const statuses = { submitted: "SUBMITTED", processing: "IN_PROGRESS", succeed: "SUCCESS", failed: "FAILURE" };
const status = statuses[data.task_status]; 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 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 || "" }; 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; if (videos.length && videos[0].url) result.url = videos[0].url;
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "OpenAI Sora video generation (text-to-video, image-to-video, and remix)", en: "OpenAI Sora video generation (text-to-video, image-to-video, and remix)",
zh: "OpenAI Sora 视频生成(文生视频、图生视频、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 channelTypes: [55, 1], // OpenAI-type channels natively serve sora with the same wire format
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
models: ["sora-2", "sora-2-pro"], models: ["sora-2", "sora-2-pro"],
...@@ -145,7 +145,9 @@ export function parseTaskResult(ctx, body) { ...@@ -145,7 +145,9 @@ export function parseTaskResult(ctx, body) {
failed: "FAILURE", failed: "FAILURE",
cancelled: "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 (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"; if (result.status === "FAILURE") result.reason = body.error && body.error.message ? body.error.message : "task failed";
return result; return result;
......
...@@ -9,7 +9,7 @@ export const meta = { ...@@ -9,7 +9,7 @@ export const meta = {
en: "SunoAPI project music and lyrics generation", en: "SunoAPI project music and lyrics generation",
zh: "SunoAPI 项目 音乐与歌词生成", zh: "SunoAPI 项目 音乐与歌词生成",
}, },
version: "1.0.0", version: "1.0.1",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [36], channelTypes: [36],
models: ["suno_music", "suno_lyrics"], models: ["suno_music", "suno_lyrics"],
...@@ -133,19 +133,23 @@ export function extractUsage(ctx) { ...@@ -133,19 +133,23 @@ export function extractUsage(ctx) {
return { clips: action === "lyrics" ? 1 : 2, action: action }; return { clips: action === "lyrics" ? 1 : 2, action: action };
} }
export function buildBatchQueryRequest(ctx, taskIds) { export function buildBatchQueryRequest(ctx, tasks) {
return { return {
url: ctx.baseUrl + "/suno/fetch", url: ctx.baseUrl + "/suno/fetch",
method: "POST", method: "POST",
headers: { "Content-Type": "application/json", Authorization: "Bearer " + ctx.apiKey }, 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 // Required v1 per-task hooks remain defined for contract compatibility. Suno's
// host polling path uses the batch hooks below. // host polling path uses the batch hooks below.
export function buildQueryRequest(ctx) { export function buildQueryRequest(ctx) {
return buildBatchQueryRequest(ctx, (ctx.requestBody || {}).ids || []); return buildBatchQueryRequest(ctx, [ctx]);
} }
export function parseBatchResult(ctx, body) { export function parseBatchResult(ctx, body) {
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -7,7 +7,7 @@ export const meta = {
en: "Google Veo video generation on Vertex AI (text-to-video and image-to-video)", en: "Google Veo video generation on Vertex AI (text-to-video and image-to-video)",
zh: "Google Veo 视频生成(文生视频、图生视频),Vertex AI 版本", zh: "Google Veo 视频生成(文生视频、图生视频),Vertex AI 版本",
}, },
version: "1.0.0", version: "1.0.1",
channelTypes: [41], channelTypes: [41],
author: { name: "QuantumNous" }, 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"], 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) { ...@@ -225,7 +225,11 @@ export function buildQueryRequest(ctx) {
} }
export function parseTaskResult(ctx, body) { export function parseTaskResult(ctx, body) {
if (body.error && body.error.message) return { status: "FAILURE", progress: "100%", reason: body.error.message }; 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 || {}); const url = dataVideo(body.response || {});
return { status: "SUCCESS", progress: "100%", url: url, remoteUrl: url }; return { status: "SUCCESS", progress: "100%", url: url, remoteUrl: url };
} }
......
...@@ -7,7 +7,7 @@ export const meta = { ...@@ -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)", en: "Shengshu Vidu video generation (text-to-video, image-to-video, first-and-last-frame, and reference-to-video)",
zh: "生数 Vidu 视频生成(文生视频、图生视频、首尾帧、参考生视频)", zh: "生数 Vidu 视频生成(文生视频、图生视频、首尾帧、参考生视频)",
}, },
version: "1.0.0", version: "1.0.1",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [52], channelTypes: [52],
models: ["viduq2", "viduq1", "vidu2.0", "vidu1.5"], models: ["viduq2", "viduq1", "vidu2.0", "vidu1.5"],
...@@ -265,7 +265,7 @@ export function buildQueryRequest(ctx) { ...@@ -265,7 +265,7 @@ export function buildQueryRequest(ctx) {
export function parseTaskResult(ctx, body) { export function parseTaskResult(ctx, body) {
const statuses = { created: "SUBMITTED", queueing: "SUBMITTED", processing: "IN_PROGRESS", success: "SUCCESS", failed: "FAILURE" }; const statuses = { created: "SUBMITTED", queueing: "SUBMITTED", processing: "IN_PROGRESS", success: "SUCCESS", failed: "FAILURE" };
const status = statuses[body.state]; 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 url = body.creations && body.creations.length ? body.creations[0].url || "" : "";
const result = { status: status, reason: body.state === "failed" ? body.err_code || "" : "" }; const result = { status: status, reason: body.state === "failed" ? body.err_code || "" : "" };
if (url) result.url = url; 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 { ...@@ -76,8 +76,8 @@ type TaskAdaptor interface {
// ── Polling ────────────────────────────────────────────────────── // ── Polling ──────────────────────────────────────────────────────
FetchTask(baseUrl, key string, body map[string]any, proxy string) (*http.Response, error) FetchTask(baseUrl, key string, task *model.Task, proxy string) (*http.Response, error)
ParseTaskResult(respBody []byte) (*relaycommon.TaskInfo, error) ParseTaskResult(task *model.Task, resp *http.Response, respBody []byte) (*relaycommon.TaskInfo, error)
} }
// TaskSubmitResponse is the transport-independent result of parsing an // TaskSubmitResponse is the transport-independent result of parsing an
...@@ -87,6 +87,7 @@ type TaskSubmitResponse struct { ...@@ -87,6 +87,7 @@ type TaskSubmitResponse struct {
TaskData []byte TaskData []byte
ClientResponse any ClientResponse any
Immediate *relaycommon.TaskInfo Immediate *relaycommon.TaskInfo
PluginState []byte
} }
type OpenAIVideoConverter interface { type OpenAIVideoConverter interface {
......
...@@ -22,12 +22,12 @@ import ( ...@@ -22,12 +22,12 @@ import (
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto" "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/logger"
"github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/model"
pluginruntime "github.com/QuantumNous/new-api/pkg/jsplugin" pluginruntime "github.com/QuantumNous/new-api/pkg/jsplugin"
"github.com/QuantumNous/new-api/relay/channel" "github.com/QuantumNous/new-api/relay/channel"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
kitdto "github.com/QuantumNous/new-api/relaykit/dto"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
...@@ -56,6 +56,7 @@ type submitResponse struct { ...@@ -56,6 +56,7 @@ type submitResponse struct {
TaskID string `json:"taskId"` TaskID string `json:"taskId"`
TaskData any `json:"taskData"` TaskData any `json:"taskData"`
Immediate *taskResult `json:"immediate"` Immediate *taskResult `json:"immediate"`
State any `json:"state"`
} }
type taskResult struct { type taskResult struct {
Code int `json:"code"` Code int `json:"code"`
...@@ -67,12 +68,16 @@ type taskResult struct { ...@@ -67,12 +68,16 @@ type taskResult struct {
RemoteURL string `json:"remoteUrl"` RemoteURL string `json:"remoteUrl"`
CompletionTokens float64 `json:"completionTokens"` CompletionTokens float64 `json:"completionTokens"`
TotalTokens float64 `json:"totalTokens"` TotalTokens float64 `json:"totalTokens"`
State any `json:"state"`
} }
var taskArtifactKeyPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._~-]{0,127}$`) var taskArtifactKeyPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._~-]{0,127}$`)
const maxTaskArtifacts = 64 const maxTaskArtifacts = 64
// maxTaskPluginPersistedJSONBytes is the shared ceiling for taskData and plugin state.
const maxTaskPluginPersistedJSONBytes = 1 << 20
type TaskAdaptor struct { type TaskAdaptor struct {
plugin *pluginruntime.LoadedPlugin plugin *pluginruntime.LoadedPlugin
info *relaycommon.RelayInfo info *relaycommon.RelayInfo
...@@ -497,10 +502,16 @@ func (a *TaskAdaptor) ParseResponse(c *gin.Context, resp *http.Response, info *r ...@@ -497,10 +502,16 @@ func (a *TaskAdaptor) ParseResponse(c *gin.Context, resp *http.Response, info *r
immediate != nil, immediate != nil,
time.Since(started).Milliseconds(), 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{ return &channel.TaskSubmitResponse{
UpstreamTaskID: parsed.TaskID, UpstreamTaskID: parsed.TaskID,
TaskData: taskData, TaskData: taskData,
Immediate: immediate, Immediate: immediate,
PluginState: pluginState,
}, nil }, nil
} }
...@@ -508,50 +519,32 @@ func (a *TaskAdaptor) GetModelList() []string { return append([]string(nil), a.p ...@@ -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) GetChannelName() string { return a.plugin.Meta.Name }
func (a *TaskAdaptor) FetchMode() string { return a.plugin.Meta.FetchMode } func (a *TaskAdaptor) FetchMode() string { return a.plugin.Meta.FetchMode }
func (a *TaskAdaptor) FetchBatchTasks(baseURL, key string, taskIDs []string, proxy string) (*http.Response, error) { func (a *TaskAdaptor) FetchBatchTasks(baseURL, key string, tasks []*model.Task, proxy string) (*http.Response, error) {
ctx := map[string]any{"baseUrl": baseURL} taskContexts := make([]map[string]any, 0, len(tasks))
auth, err := resolveAuth(a.plugin.Meta.Auth, key, proxy) 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 { if err != nil {
return nil, err return nil, err
} }
ctx["auth"] = auth value, err := a.plugin.Engine.Call(context.Background(), "buildBatchQueryRequest", ctx, taskContexts)
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)
if err != nil { if err != nil {
return nil, err return nil, err
} }
return a.doFetchDescriptor(baseURL, proxy, value) return a.doFetchDescriptor(baseURL, proxy, value)
} }
func (a *TaskAdaptor) FetchTask(baseURL, key string, body map[string]any, proxy string) (*http.Response, error) { func (a *TaskAdaptor) FetchTask(baseURL, key string, task *model.Task, proxy string) (*http.Response, error) {
ctx := map[string]any{"taskId": body["task_id"], "action": body["action"], "requestBody": body, "baseUrl": baseURL} ctx, err := a.queryContext(task, key, baseURL, proxy)
// 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)
if err != nil { if err != nil {
return nil, err return nil, err
} }
ctx["auth"] = auth value, err := a.plugin.Engine.Call(context.Background(), "buildQueryRequest", ctx)
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)
if err != nil { if err != nil {
return nil, err return nil, err
} }
...@@ -616,14 +609,27 @@ func (a *TaskAdaptor) doFetchDescriptor(baseURL, proxy string, value any) (*http ...@@ -616,14 +609,27 @@ func (a *TaskAdaptor) doFetchDescriptor(baseURL, proxy string, value any) (*http
return resp, nil 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() started := time.Now()
input := any(string(body)) input := any(string(body))
var decoded any var decoded any
if common.Unmarshal(body, &decoded) == nil { if common.Unmarshal(body, &decoded) == nil {
input = decoded 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 { 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()) 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 return nil, err
...@@ -639,6 +645,7 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa ...@@ -639,6 +645,7 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa
StartTime int64 `json:"startTime"` StartTime int64 `json:"startTime"`
FinishTime int64 `json:"finishTime"` FinishTime int64 `json:"finishTime"`
Data any `json:"data"` Data any `json:"data"`
State any `json:"state"`
} }
if err = convert(value, &parsed); err != nil { 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()) 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 ...@@ -651,12 +658,27 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa
continue continue
} }
info := relaycommon.TaskInfo{TaskID: item.TaskID, Status: item.Status, Progress: item.Progress, Reason: item.Reason, Url: item.URL} 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 { if hasCompletionUsage {
usageBody := item.Data usageBody := item.Data
if usageBody == nil { if usageBody == nil {
usageBody = jsonValue(item) 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 { if hookErr == nil {
a.applyCompletionUsageFacts(&info, facts) a.applyCompletionUsageFacts(&info, facts)
} }
...@@ -675,14 +697,22 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa ...@@ -675,14 +697,22 @@ func (a *TaskAdaptor) ParseBatchResult(body []byte) (map[string]*service.BatchTa
return results, nil 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() started := time.Now()
input := any(string(body)) input := any(string(body))
var decoded any var decoded any
if common.Unmarshal(body, &decoded) == nil { if common.Unmarshal(body, &decoded) == nil {
input = decoded 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 { 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()) 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 return nil, err
...@@ -703,10 +733,17 @@ func (a *TaskAdaptor) ParseTaskResult(body []byte) (*relaycommon.TaskInfo, error ...@@ -703,10 +733,17 @@ func (a *TaskAdaptor) ParseTaskResult(body []byte) (*relaycommon.TaskInfo, error
CompletionTokens: positiveInt(parsed.CompletionTokens), CompletionTokens: positiveInt(parsed.CompletionTokens),
TotalTokens: positiveInt(parsed.TotalTokens), 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 // The raw polling response only exists at this boundary. Capture upstream
// units here so the host settlement path can consume them from TaskInfo. // units here so the host settlement path can consume them from TaskInfo.
if a.hasHook(context.Background(), "extractUsageOnComplete") { 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 { if hookErr == nil {
a.applyCompletionUsageFacts(result, facts) a.applyCompletionUsageFacts(result, facts)
} }
...@@ -901,15 +938,125 @@ func taskArtifactContext(task *model.Task) (map[string]any, error) { ...@@ -901,15 +938,125 @@ func taskArtifactContext(task *model.Task) (map[string]any, error) {
if task.PrivateData.Execution != nil && task.PrivateData.Execution.TaskPlugin != nil { if task.PrivateData.Execution != nil && task.PrivateData.Execution.TaskPlugin != nil {
producerVersion = task.PrivateData.Execution.TaskPlugin.Version 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{ return map[string]any{
"taskId": task.TaskID, "taskId": task.TaskID,
"status": string(task.Status), "status": string(task.Status),
"action": task.Action, "action": task.Action,
"data": data, "data": data,
"state": state,
"producerVersion": producerVersion, "producerVersion": producerVersion,
}, nil }, 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) { func validateTaskArtifacts(value any) ([]channel.TaskArtifact, error) {
encoded, err := common.Marshal(value) encoded, err := common.Marshal(value)
if err != nil { if err != nil {
......
...@@ -503,12 +503,15 @@ func TestTaskAdaptorMapsJSContract(t *testing.T) { ...@@ -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.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}`))) 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) require.NoError(t, err)
queryBody, err := io.ReadAll(queryResp.Body) queryBody, err := io.ReadAll(queryResp.Body)
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, queryResp.Body.Close()) require.NoError(t, queryResp.Body.Close())
result, err := adaptor.ParseTaskResult(queryBody) result, err := adaptor.ParseTaskResult(&model.Task{}, queryResp, queryBody)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "SUCCESS", result.Status) assert.Equal(t, "SUCCESS", result.Status)
assert.Equal(t, "https://cdn.example/video.mp4", result.Url) assert.Equal(t, "https://cdn.example/video.mp4", result.Url)
...@@ -866,7 +869,7 @@ export function extractUsageOnComplete(task, result, body) { return (body || {}) ...@@ -866,7 +869,7 @@ export function extractUsageOnComplete(task, result, body) { return (body || {})
adaptor, _, _ := newRequest(t, map[string]any{}) adaptor, _, _ := newRequest(t, map[string]any{})
body, marshalErr := common.Marshal(map[string]any{"completionUsage": testCase.usage}) body, marshalErr := common.Marshal(map[string]any{"completionUsage": testCase.usage})
require.NoError(t, marshalErr) 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) require.NoError(t, parseErr)
assert.Nil(t, result.UsageFacts) assert.Nil(t, result.UsageFacts)
assert.Zero(t, result.TotalTokens) assert.Zero(t, result.TotalTokens)
...@@ -877,7 +880,7 @@ export function extractUsageOnComplete(task, result, body) { return (body || {}) ...@@ -877,7 +880,7 @@ export function extractUsageOnComplete(task, result, body) { return (body || {})
adaptor, _, _ := newRequest(t, map[string]any{}) adaptor, _, _ := newRequest(t, map[string]any{})
body, err := common.Marshal(map[string]any{"completionUsage": map[string]any{"tokens": 500000}}) body, err := common.Marshal(map[string]any{"completionUsage": map[string]any{"tokens": 500000}})
require.NoError(t, err) 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) require.NoError(t, err)
assert.EqualValues(t, 500000, result.UsageFacts["tokens"]) assert.EqualValues(t, 500000, result.UsageFacts["tokens"])
}) })
...@@ -916,7 +919,7 @@ export function extractUsageOnComplete() { return {units: 3.5}; } ...@@ -916,7 +919,7 @@ export function extractUsageOnComplete() { return {units: 3.5}; }
body, err := common.Marshal(map[string]any{}) body, err := common.Marshal(map[string]any{})
require.NoError(t, err) 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) require.NoError(t, err)
assert.Equal(t, 3.5, result.UsageFacts["units"]) assert.Equal(t, 3.5, result.UsageFacts["units"])
}) })
...@@ -925,7 +928,7 @@ export function extractUsageOnComplete() { return {units: 3.5}; } ...@@ -925,7 +928,7 @@ export function extractUsageOnComplete() { return {units: 3.5}; }
adaptor, _, _ := newRequest(t, map[string]any{}) adaptor, _, _ := newRequest(t, map[string]any{})
body, err := common.Marshal(map[string]any{"completionUsage": map[string]any{"upstreamUnits": 5000}}) body, err := common.Marshal(map[string]any{"completionUsage": map[string]any{"upstreamUnits": 5000}})
require.NoError(t, err) 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) require.NoError(t, err)
assert.Equal(t, 5000, result.TotalTokens) assert.Equal(t, 5000, result.TotalTokens)
assert.EqualValues(t, 5000, result.UsageFacts["upstreamUnits"]) assert.EqualValues(t, 5000, result.UsageFacts["upstreamUnits"])
...@@ -988,7 +991,7 @@ export function parseTaskResult(ctx, body) { return {status: "SUCCESS", completi ...@@ -988,7 +991,7 @@ export function parseTaskResult(ctx, body) { return {status: "SUCCESS", completi
require.NoError(t, err) require.NoError(t, err)
adaptor := New(plugin) 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) require.NoError(t, err)
assert.Equal(t, 13, result.CompletionTokens) assert.Equal(t, 13, result.CompletionTokens)
assert.Equal(t, 17, result.TotalTokens) assert.Equal(t, 17, result.TotalTokens)
...@@ -1093,7 +1096,7 @@ export function buildSubmitRequest(ctx) { return { url: ctx.baseUrl + "/submit", ...@@ -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 parseSubmitResponse(ctx, resp) { return { taskId: resp.body.id }; }
export function buildQueryRequest(ctx) { return { url: ctx.baseUrl + "/tasks/" + ctx.taskId }; } export function buildQueryRequest(ctx) { return { url: ctx.baseUrl + "/tasks/" + ctx.taskId }; }
export function parseTaskResult(ctx, body) { return { taskId: body.id, status: "SUCCESS" }; } 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) { export function parseBatchResult(ctx, body) {
return body.items.map(function (item) { 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 }; 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) { ...@@ -1128,13 +1131,17 @@ func TestTaskAdaptorBatchBridge(t *testing.T) {
adaptor := New(plugin) adaptor := New(plugin)
require.Equal(t, "batch", adaptor.FetchMode()) 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) require.NoError(t, err)
defer resp.Body.Close() defer resp.Body.Close()
payload, err := io.ReadAll(resp.Body) payload, err := io.ReadAll(resp.Body)
require.NoError(t, err) require.NoError(t, err)
results, err := adaptor.ParseBatchResult(payload) results, err := adaptor.ParseBatchResult(tasks, resp, payload)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, results, 2, "entry without taskId must be skipped") require.Len(t, results, 2, "entry without taskId must be skipped")
...@@ -1233,26 +1240,176 @@ export function parseTaskResult(){return {status:"SUCCESS"}} ...@@ -1233,26 +1240,176 @@ export function parseTaskResult(){return {status:"SUCCESS"}}
testCases := []struct { testCases := []struct {
name string name string
body map[string]any task *model.Task
want string want string
}{ }{
{ {
name: "mapped model", 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", want: "/tasks/alias/declared-model/t1",
}, },
{ {
name: "unmapped model falls back to the origin name", 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", want: "/tasks/alias/alias/t1",
}, },
} }
for _, testCase := range testCases { for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) { 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, fetchErr)
require.NoError(t, resp.Body.Close()) require.NoError(t, resp.Body.Close())
assert.Equal(t, testCase.want, requested) 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 { ...@@ -119,19 +119,19 @@ type RelayInfo struct {
ReasoningEffort string ReasoningEffort string
// ReasoningConversion is the suffix-derived reasoning intent attached // ReasoningConversion is the suffix-derived reasoning intent attached
// after model mapping. Converters read it via ReasoningState(). // after model mapping. Converters read it via ReasoningState().
ReasoningConversion *dto.ReasoningConversionState ReasoningConversion *dto.ReasoningConversionState
UserSetting dto.UserSetting UserSetting dto.UserSetting
UserEmail string UserEmail string
UserQuota int UserQuota int
RelayFormat types.RelayFormat RelayFormat types.RelayFormat
SendResponseCount int SendResponseCount int
// ClaudeToChatStreamState / ChatToGeminiStreamState hold per-attempt // ClaudeToChatStreamState / ChatToGeminiStreamState hold per-attempt
// stream converters. InitChannelMeta nils them so a retry cannot resume a // stream converters. InitChannelMeta nils them so a retry cannot resume a
// dirty converter (advanced tool index / finalized). // dirty converter (advanced tool index / finalized).
ClaudeToChatStreamState any ClaudeToChatStreamState any
ChatToGeminiStreamState any ChatToGeminiStreamState any
ReceivedResponseCount int ReceivedResponseCount int
FinalPreConsumedQuota int // 最终预消耗的配额 FinalPreConsumedQuota int // 最终预消耗的配额
// ForcePreConsume 为 true 时禁用 BillingSession 的信任额度旁路, // ForcePreConsume 为 true 时禁用 BillingSession 的信任额度旁路,
// 强制预扣全额。用于异步任务(视频/音乐生成等),因为请求返回后任务仍在运行, // 强制预扣全额。用于异步任务(视频/音乐生成等),因为请求返回后任务仍在运行,
// 必须在提交前锁定全额。 // 必须在提交前锁定全额。
...@@ -1018,16 +1018,17 @@ func (t *TaskSubmitReq) UnmarshalMetadata(v any) error { ...@@ -1018,16 +1018,17 @@ func (t *TaskSubmitReq) UnmarshalMetadata(v any) error {
} }
type TaskInfo struct { type TaskInfo struct {
Code int `json:"code"` Code int `json:"code"`
TaskID string `json:"task_id"` TaskID string `json:"task_id"`
Status string `json:"status"` Status string `json:"status"`
Reason string `json:"reason,omitempty"` Reason string `json:"reason,omitempty"`
Url string `json:"url,omitempty"` Url string `json:"url,omitempty"`
RemoteUrl string `json:"remote_url,omitempty"` RemoteUrl string `json:"remote_url,omitempty"`
Progress string `json:"progress,omitempty"` Progress string `json:"progress,omitempty"`
CompletionTokens int `json:"completion_tokens,omitempty"` // 用于按倍率计费 CompletionTokens int `json:"completion_tokens,omitempty"` // 用于按倍率计费
TotalTokens int `json:"total_tokens,omitempty"` // 用于按倍率计费 TotalTokens int `json:"total_tokens,omitempty"` // 用于按倍率计费
UsageFacts map[string]any `json:"usage_facts,omitempty"` UsageFacts map[string]any `json:"usage_facts,omitempty"`
PluginState json.RawMessage `json:"plugin_state,omitempty"`
} }
func FailTaskInfo(reason string) *TaskInfo { func FailTaskInfo(reason string) *TaskInfo {
......
...@@ -32,6 +32,7 @@ type TaskSubmitResult struct { ...@@ -32,6 +32,7 @@ type TaskSubmitResult struct {
Platform constant.TaskPlatform Platform constant.TaskPlatform
Quota int Quota int
Immediate *relaycommon.TaskInfo Immediate *relaycommon.TaskInfo
PluginState []byte
//PerCallPrice types.PriceData //PerCallPrice types.PriceData
} }
...@@ -381,6 +382,7 @@ func RelayTaskSubmit(c *gin.Context, info *relaycommon.RelayInfo) (*TaskSubmitRe ...@@ -381,6 +382,7 @@ func RelayTaskSubmit(c *gin.Context, info *relaycommon.RelayInfo) (*TaskSubmitRe
Platform: platform, Platform: platform,
Quota: finalQuota, Quota: finalQuota,
Immediate: parsed.Immediate, Immediate: parsed.Immediate,
PluginState: parsed.PluginState,
}, nil }, nil
} }
...@@ -517,12 +519,7 @@ func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte { ...@@ -517,12 +519,7 @@ func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte {
return nil return nil
} }
resp, err := adaptor.FetchTask(baseURL, channelModel.Key, map[string]any{ resp, err := adaptor.FetchTask(baseURL, channelModel.Key, task, proxy)
"task_id": task.GetUpstreamTaskID(),
"action": constant.NormalizeTaskAction(task.Action),
"model": task.Properties.OriginModelName,
"upstream_model": task.Properties.UpstreamModelName,
}, proxy)
if err != nil || resp == nil { if err != nil || resp == nil {
return nil return nil
} }
...@@ -532,7 +529,7 @@ func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte { ...@@ -532,7 +529,7 @@ func tryRealtimeFetch(task *model.Task, isOpenAIVideoAPI bool) []byte {
return nil return nil
} }
ti, err := adaptor.ParseTaskResult(body) ti, err := adaptor.ParseTaskResult(task, resp, body)
if err != nil || ti == nil { if err != nil || ti == nil {
return nil return nil
} }
......
...@@ -1343,10 +1343,12 @@ type mockAdaptor struct { ...@@ -1343,10 +1343,12 @@ type mockAdaptor struct {
} }
func (m *mockAdaptor) Init(_ *relaycommon.RelayInfo) {} 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 return nil, nil
} }
func (m *mockAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) { return nil, nil }
func (m *mockAdaptor) AdjustBillingOnComplete(_ *model.Task, _ *relaycommon.TaskInfo) int { func (m *mockAdaptor) AdjustBillingOnComplete(_ *model.Task, _ *relaycommon.TaskInfo) int {
return m.adjustReturn return m.adjustReturn
} }
......
...@@ -62,3 +62,24 @@ func TestBuildTaskPluginViewRewritesOnlyStructuredTaskIDFields(t *testing.T) { ...@@ -62,3 +62,24 @@ func TestBuildTaskPluginViewRewritesOnlyStructuredTaskIDFields(t *testing.T) {
assert.Equal(t, privateTaskID, nested[1]) 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 ( ...@@ -18,7 +18,6 @@ import (
"github.com/QuantumNous/new-api/pkg/billingexpr" "github.com/QuantumNous/new-api/pkg/billingexpr"
"github.com/QuantumNous/new-api/relay/channel/task/taskcommon" "github.com/QuantumNous/new-api/relay/channel/task/taskcommon"
relaycommon "github.com/QuantumNous/new-api/relay/common" relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relaykit/dto"
"github.com/bytedance/gopkg/util/gopool" "github.com/bytedance/gopkg/util/gopool"
"github.com/samber/lo" "github.com/samber/lo"
...@@ -27,8 +26,8 @@ import ( ...@@ -27,8 +26,8 @@ import (
// TaskPollingAdaptor 定义轮询所需的最小适配器接口,避免 service -> relay 的循环依赖 // TaskPollingAdaptor 定义轮询所需的最小适配器接口,避免 service -> relay 的循环依赖
type TaskPollingAdaptor interface { type TaskPollingAdaptor interface {
Init(info *relaycommon.RelayInfo) Init(info *relaycommon.RelayInfo)
FetchTask(baseURL string, key string, body map[string]any, proxy string) (*http.Response, error) FetchTask(baseURL string, key string, task *model.Task, proxy string) (*http.Response, error)
ParseTaskResult(body []byte) (*relaycommon.TaskInfo, error) ParseTaskResult(task *model.Task, resp *http.Response, body []byte) (*relaycommon.TaskInfo, error)
// AdjustBillingOnComplete 在任务到达终态(成功/失败)时由轮询循环调用。 // AdjustBillingOnComplete 在任务到达终态(成功/失败)时由轮询循环调用。
// 返回正数触发差额结算(补扣/退还),返回 0 保持预扣费金额不变。 // 返回正数触发差额结算(补扣/退还),返回 0 保持预扣费金额不变。
AdjustBillingOnComplete(task *model.Task, taskResult *relaycommon.TaskInfo) int AdjustBillingOnComplete(task *model.Task, taskResult *relaycommon.TaskInfo) int
...@@ -37,10 +36,21 @@ type TaskPollingAdaptor interface { ...@@ -37,10 +36,21 @@ type TaskPollingAdaptor interface {
type BatchTaskPollingAdaptor interface { type BatchTaskPollingAdaptor interface {
TaskPollingAdaptor TaskPollingAdaptor
FetchMode() string FetchMode() string
FetchBatchTasks(baseURL, key string, taskIDs []string, proxy string) (*http.Response, error) FetchBatchTasks(baseURL, key string, tasks []*model.Task, proxy string) (*http.Response, error)
ParseBatchResult(body []byte) (map[string]*BatchTaskResult, 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 { type BatchTaskResult struct {
TaskInfo relaycommon.TaskInfo TaskInfo relaycommon.TaskInfo
Action string Action string
...@@ -258,24 +268,39 @@ func updateBatchTasks(ctx context.Context, adaptor BatchTaskPollingAdaptor, chan ...@@ -258,24 +268,39 @@ func updateBatchTasks(ctx context.Context, adaptor BatchTaskPollingAdaptor, chan
if baseURL == "" { if baseURL == "" {
baseURL = constant.GetChannelBaseURL(ch.Type) 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 { if err != nil {
common.SysLog(fmt.Sprintf("Get Task Do req error: %v", err)) common.SysLog(fmt.Sprintf("Get Task Do req error: %v", err))
return err return recordPollFailureForTasks(ctx, adaptor, tasks, pollClassTransport, 0, err.Error())
}
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)
} }
defer resp.Body.Close() defer resp.Body.Close()
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { if err != nil {
common.SysLog(fmt.Sprintf("Get Suno Task parse body error: %v", err)) common.SysLog(fmt.Sprintf("Get Suno Task parse body error: %v", err))
return err return recordPollFailureForTasks(ctx, adaptor, tasks, pollClassTransport, resp.StatusCode, err.Error())
} }
responseItems, err := adaptor.ParseBatchResult(responseBody) 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 { 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 { for upstreamID, responseItem := range responseItems {
if ctx.Err() != nil { if ctx.Err() != nil {
...@@ -287,7 +312,27 @@ func updateBatchTasks(ctx context.Context, adaptor BatchTaskPollingAdaptor, chan ...@@ -287,7 +312,27 @@ func updateBatchTasks(ctx context.Context, adaptor BatchTaskPollingAdaptor, chan
continue continue
} }
snap := task.Snapshot() 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.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.SubmitTime = lo.If(responseItem.SubmitTime != 0, responseItem.SubmitTime).Else(task.SubmitTime)
task.StartTime = lo.If(responseItem.StartTime != 0, responseItem.StartTime).Else(task.StartTime) task.StartTime = lo.If(responseItem.StartTime != 0, responseItem.StartTime).Else(task.StartTime)
...@@ -447,24 +492,28 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch * ...@@ -447,24 +492,28 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch *
if privateData.Key != "" { if privateData.Key != "" {
key = privateData.Key key = privateData.Key
} }
resp, err := adaptor.FetchTask(baseURL, key, map[string]any{ snap := task.Snapshot()
"task_id": task.GetUpstreamTaskID(), resp, err := adaptor.FetchTask(baseURL, key, task, proxy)
"action": constant.NormalizeTaskAction(task.Action),
"model": task.Properties.OriginModelName,
"upstream_model": task.Properties.UpstreamModelName,
}, proxy)
if err != nil { 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() defer resp.Body.Close()
responseBody, err := io.ReadAll(resp.Body) responseBody, err := io.ReadAll(resp.Body)
if err != nil { 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) 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{} taskResult := &relaycommon.TaskInfo{}
// try parse as New API response format // try parse as New API response format
...@@ -478,41 +527,33 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch * ...@@ -478,41 +527,33 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch *
taskResult.Progress = t.Progress taskResult.Progress = t.Progress
taskResult.Reason = t.FailReason taskResult.Reason = t.FailReason
task.Data = t.Data task.Data = t.Data
} else if taskResult, err = adaptor.ParseTaskResult(responseBody); err != nil { } else if taskResult, err = adaptor.ParseTaskResult(task, resp, responseBody); err != nil {
return fmt.Errorf("parseTaskResult failed for task %s: %w", taskId, err) return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassHookError, resp.StatusCode, err.Error())
} }
task.Data = redactVideoResponseBody(responseBody)
logger.LogDebug(ctx, "updateVideoSingleTask taskResult: %+v", taskResult) logger.LogDebug(ctx, "updateVideoSingleTask taskResult: %+v", taskResult)
now := time.Now().Unix() parsedStatus := model.TaskStatus(taskResult.Status)
if taskResult.Status == "" { if parsedStatus == model.TaskStatusUnknown || parsedStatus == "" || !knownPollStatus(parsedStatus) {
//taskResult = relaycommon.FailTaskInfo("upstream returned empty status") return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassUnrecognized, resp.StatusCode, unrecognizedPollDetail(taskResult.Reason, responseBody))
errorResult := &dto.GeneralErrorResponse{} }
if err = common.Unmarshal(responseBody, &errorResult); err == nil { if classifyPollHTTP(resp.StatusCode) == pollClassOtherClient && isNonTerminalPollStatus(parsedStatus) {
openaiError := errorResult.TryToOpenAIError() return recordPollFailure(ctx, adaptor, task, snap.Status, pollClassUnrecognized, resp.StatusCode, unrecognizedPollDetail(taskResult.Reason, responseBody))
if openaiError != nil { }
// 返回规范的 OpenAI 错误格式,提取错误信息,判断错误是否为任务失败
if openaiError.Code == "429" {
// 429 错误通常表示请求过多或速率限制,暂时不认为是任务失败,保持原状态等待下一轮轮询
return nil
}
// 其他错误认为是任务失败,记录错误信息并更新任务状态 task.Data = redactVideoResponseBody(responseBody)
taskResult = relaycommon.FailTaskInfo("upstream returned error") if len(taskResult.PluginState) > 0 {
} else { task.PrivateData.PluginState = taskResult.PluginState
// 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))) if isNonTerminalPollStatus(parsedStatus) {
taskResult = relaycommon.FailTaskInfo("upstream returned unrecognized message") task.PrivateData.PollFailures = 0
}
}
} }
now := time.Now().Unix()
shouldFinalizeBilling := false shouldFinalizeBilling := false
task.Status = model.TaskStatus(taskResult.Status) task.Status = parsedStatus
switch taskResult.Status { switch parsedStatus {
case model.TaskStatusSubmitted: case model.TaskStatusSubmitted:
task.Progress = taskcommon.ProgressSubmitted task.Progress = taskcommon.ProgressSubmitted
case model.TaskStatusQueued: case model.TaskStatusQueued:
...@@ -549,8 +590,6 @@ func updateVideoSingleTask(ctx context.Context, adaptor TaskPollingAdaptor, ch * ...@@ -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)) logger.LogInfo(ctx, fmt.Sprintf("Task %s failed: %s", task.TaskID, task.FailReason))
taskResult.Progress = taskcommon.ProgressComplete taskResult.Progress = taskcommon.ProgressComplete
shouldFinalizeBilling = true shouldFinalizeBilling = true
default:
return fmt.Errorf("unknown task status %s for task %s", taskResult.Status, task.TaskID)
} }
if taskResult.Progress != "" { if taskResult.Progress != "" {
task.Progress = taskResult.Progress task.Progress = taskResult.Progress
...@@ -670,3 +709,131 @@ func settleTaskBillingOnComplete(ctx context.Context, adaptor TaskPollingAdaptor ...@@ -670,3 +709,131 @@ func settleTaskBillingOnComplete(ctx context.Context, adaptor TaskPollingAdaptor
} }
return false 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 { ...@@ -40,12 +40,15 @@ type batchPollingAdaptor struct {
} }
func (a *batchPollingAdaptor) FetchMode() string { return "batch" } 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.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 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 { if a.results != nil {
return a.results, nil return a.results, nil
} }
...@@ -58,8 +61,11 @@ func (a *batchPollingAdaptor) ParseBatchResult([]byte) (map[string]*BatchTaskRes ...@@ -58,8 +61,11 @@ func (a *batchPollingAdaptor) ParseBatchResult([]byte) (map[string]*BatchTaskRes
func (a *taskPollingFetchAdaptor) Init(_ *relaycommon.RelayInfo) {} func (a *taskPollingFetchAdaptor) Init(_ *relaycommon.RelayInfo) {}
func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, body map[string]any, _ string) (*http.Response, error) { func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, task *model.Task, _ string) (*http.Response, error) {
taskID, _ := body["task_id"].(string) taskID := ""
if task != nil {
taskID = task.GetUpstreamTaskID()
}
if taskID == a.blockTaskID && a.releaseBlock != nil { if taskID == a.blockTaskID && a.releaseBlock != nil {
a.blockOnce.Do(func() { a.blockOnce.Do(func() {
if a.blockStarted != nil { if a.blockStarted != nil {
...@@ -97,7 +103,7 @@ func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, body map[string] ...@@ -97,7 +103,7 @@ func (a *taskPollingFetchAdaptor) FetchTask(_ string, _ string, body map[string]
}, nil }, 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 return &relaycommon.TaskInfo{Status: model.TaskStatusInProgress}, nil
} }
...@@ -721,3 +727,273 @@ func TestSweepTimedOutTasksHonorsRefundRolloutBoundary(t *testing.T) { ...@@ -721,3 +727,273 @@ func TestSweepTimedOutTasksHonorsRefundRolloutBoundary(t *testing.T) {
assert.Equal(t, initialQuota+modernTaskQuota, getUserQuota(t, userID)) assert.Equal(t, initialQuota+modernTaskQuota, getUserQuota(t, userID))
assert.Equal(t, int64(1), countLogs(t)) 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