Commit 92bc7ff7 by CaIon

fix(plugins): make sunoapi alias-safe and lock alias echo across built-ins

Decode on ctx.upstreamModel || ctx.model and echo ctx.model; fix the
lyrics/music render branch that read a nonexistent ctx.requestBody.model;
stop sending empty Accept/Content-Type. Bump sunoapi to 1.0.2. Add an
alias-echo table test covering every built-in.
parent 6298b0f3
...@@ -4,11 +4,14 @@ import ( ...@@ -4,11 +4,14 @@ import (
"io/fs" "io/fs"
"testing" "testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/pkg/jsplugin" "github.com/QuantumNous/new-api/pkg/jsplugin"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
var expectedKeys = []string{"alibaba", "doubao", "google", "hailuo", "jimeng", "kling", "sora", "sunoapi", "vertex-ai", "vidu"}
func TestBuiltInVendorPluginsDeclareNativeRoutesAndLegacyChannelTypes(t *testing.T) { func TestBuiltInVendorPluginsDeclareNativeRoutesAndLegacyChannelTypes(t *testing.T) {
generation := jsplugin.DefaultRegistry.Generation() generation := jsplugin.DefaultRegistry.Generation()
require.NotNil(t, generation) require.NotNil(t, generation)
...@@ -63,7 +66,6 @@ func TestBuiltInVendorPluginsDeclareNativeRoutesAndLegacyChannelTypes(t *testing ...@@ -63,7 +66,6 @@ func TestBuiltInVendorPluginsDeclareNativeRoutesAndLegacyChannelTypes(t *testing
} }
func TestBuiltInTaskPluginResponsesAndUsageContracts(t *testing.T) { func TestBuiltInTaskPluginResponsesAndUsageContracts(t *testing.T) {
expectedKeys := []string{"alibaba", "doubao", "google", "hailuo", "jimeng", "kling", "sora", "sunoapi", "vertex-ai", "vidu"}
generation := jsplugin.DefaultRegistry.Generation() generation := jsplugin.DefaultRegistry.Generation()
require.NotNil(t, generation) require.NotNil(t, generation)
...@@ -121,3 +123,35 @@ func TestBuiltInTaskPluginResponsesAndUsageContracts(t *testing.T) { ...@@ -121,3 +123,35 @@ func TestBuiltInTaskPluginResponsesAndUsageContracts(t *testing.T) {
}) })
} }
} }
func TestBuiltInResponsesDecodersEchoChannelMappedAlias(t *testing.T) {
bodyOverrides := map[string]map[string]any{}
for _, key := range expectedKeys {
t.Run(key, func(t *testing.T) {
source, sourceErr := Source(key)
require.NoError(t, sourceErr)
registry := jsplugin.NewRegistry()
plugin, registerErr := registry.RegisterFactory(source, jsplugin.Options{Key: key})
require.NoError(t, registerErr)
require.NotEmpty(t, plugin.Meta.Models)
alias := "alias-under-test"
upstreamModel := plugin.Meta.Models[0]
body := map[string]any{"model": alias, "input": "a cat walking on the beach"}
if override, ok := bodyOverrides[key]; ok {
body = override
}
value, callErr := plugin.Engine.CallPath(t.Context(), "protocols", []string{"openai_responses", "decodeRequest"}, map[string]any{
"model": alias, "upstreamModel": upstreamModel, "stream": false,
"body": map[string]any{"kind": "json", "value": body},
})
require.NoError(t, callErr)
encoded, marshalErr := common.Marshal(value)
require.NoError(t, marshalErr)
var decoded map[string]any
require.NoError(t, common.Unmarshal(encoded, &decoded))
assert.Equal(t, "submit", decoded["kind"])
assert.Equal(t, alias, decoded["model"])
})
}
}
...@@ -74,6 +74,30 @@ func TestSunoResponsesProtocol(t *testing.T) { ...@@ -74,6 +74,30 @@ func TestSunoResponsesProtocol(t *testing.T) {
require.ErrorContains(t, callErr, "input must be a string or array") require.ErrorContains(t, callErr, "input must be a string or array")
}) })
t.Run("echoes a channel-mapped alias", func(t *testing.T) {
music := decodeMap(t, callProtocol(t, "decodeRequest", map[string]any{
"model": "my-suno", "upstreamModel": "suno_music", "stream": false,
"body": map[string]any{"kind": "json", "value": map[string]any{"model": "my-suno", "input": "summer pop"}},
}))
assert.Equal(t, "my-suno", music["model"])
assert.Equal(t, "MUSIC", music["action"])
lyrics := decodeMap(t, callProtocol(t, "decodeRequest", map[string]any{
"model": "my-suno", "upstreamModel": "suno_lyrics", "stream": false,
"body": map[string]any{"kind": "json", "value": map[string]any{"model": "my-suno", "input": "write about the sea"}},
}))
assert.Equal(t, "my-suno", lyrics["model"])
assert.Equal(t, "LYRICS", lyrics["action"])
assert.Equal(t, map[string]any{"prompt": "write about the sea"}, lyrics["requestBody"])
})
t.Run("rejects a model this plugin does not serve", func(t *testing.T) {
_, callErr := plugin.Engine.CallPath(t.Context(), "protocols", []string{"openai_responses", "decodeRequest"}, map[string]any{
"model": "gpt-4o", "body": map[string]any{"kind": "json", "value": map[string]any{"model": "gpt-4o", "input": "hello"}},
})
require.ErrorContains(t, callErr, "suno_music or suno_lyrics")
})
t.Run("extracts schema-declared usage", func(t *testing.T) { t.Run("extracts schema-declared usage", func(t *testing.T) {
value, callErr := plugin.Engine.Call(t.Context(), "extractUsage", map[string]any{ value, callErr := plugin.Engine.Call(t.Context(), "extractUsage", map[string]any{
"model": "suno_music", "action": "MUSIC", "usagePurpose": "facts", "requestBody": map[string]any{}, "model": "suno_music", "action": "MUSIC", "usagePurpose": "facts", "requestBody": map[string]any{},
...@@ -92,8 +116,8 @@ func TestSunoResponsesProtocol(t *testing.T) { ...@@ -92,8 +116,8 @@ func TestSunoResponsesProtocol(t *testing.T) {
const firstAudioKey = "audio-5fc0a0cd3367274b4b6de056fc754263f8726a704bb4814ffeb88495f22dad35" const firstAudioKey = "audio-5fc0a0cd3367274b4b6de056fc754263f8726a704bb4814ffeb88495f22dad35"
const secondAudioKey = "audio-a9c04b840373f4ef4e8d80140b745c6f647819fa375bc34368cdccced7e2b455" const secondAudioKey = "audio-a9c04b840373f4ef4e8d80140b745c6f647819fa375bc34368cdccced7e2b455"
protocolContext := map[string]any{ protocolContext := map[string]any{
"requestBody": map[string]any{"model": "suno_music"}, "model": "suno_music",
"stream": true, "stream": true,
"artifacts": map[string]any{ "artifacts": map[string]any{
firstAudioKey: map[string]any{"key": firstAudioKey, "type": "audio", "url": "https://gateway.example/artifacts/song-1"}, firstAudioKey: map[string]any{"key": firstAudioKey, "type": "audio", "url": "https://gateway.example/artifacts/song-1"},
secondAudioKey: map[string]any{"key": secondAudioKey, "type": "audio", "url": "https://gateway.example/artifacts/song-2"}, secondAudioKey: map[string]any{"key": secondAudioKey, "type": "audio", "url": "https://gateway.example/artifacts/song-2"},
...@@ -167,7 +191,7 @@ func TestSunoResponsesProtocol(t *testing.T) { ...@@ -167,7 +191,7 @@ func TestSunoResponsesProtocol(t *testing.T) {
}) })
t.Run("renders lyrics without audio artifacts", func(t *testing.T) { t.Run("renders lyrics without audio artifacts", func(t *testing.T) {
value := callProtocol(t, "renderFinal", map[string]any{"requestBody": map[string]any{"model": "suno_lyrics"}}, map[string]any{ value := callProtocol(t, "renderFinal", map[string]any{"model": "suno_lyrics"}, map[string]any{
"task_id": "lyrics-public", "status": "SUCCESS", "data": map[string]any{"id": "lyrics-1", "title": "Tide", "text": "Sea lyrics"}, "task_id": "lyrics-public", "status": "SUCCESS", "data": map[string]any{"id": "lyrics-1", "title": "Tide", "text": "Sea lyrics"},
}) })
machine := relay.NewPluginResponsesMachine("lyrics-public", "suno_lyrics", 10, relay.DefaultPluginProtocolLimits()) machine := relay.NewPluginResponsesMachine("lyrics-public", "suno_lyrics", 10, relay.DefaultPluginProtocolLimits())
...@@ -187,9 +211,37 @@ func TestSunoResponsesProtocol(t *testing.T) { ...@@ -187,9 +211,37 @@ func TestSunoResponsesProtocol(t *testing.T) {
t.Run("requires host audio artifacts", func(t *testing.T) { t.Run("requires host audio artifacts", func(t *testing.T) {
_, callErr := plugin.Engine.CallPath(t.Context(), "protocols", []string{"openai_responses", "renderFinal"}, map[string]any{ _, callErr := plugin.Engine.CallPath(t.Context(), "protocols", []string{"openai_responses", "renderFinal"}, map[string]any{
"requestBody": map[string]any{"model": "suno_music"}, "model": "suno_music",
}, successTask) }, successTask)
require.ErrorContains(t, callErr, "audio artifact is unavailable") require.ErrorContains(t, callErr, "audio artifact is unavailable")
assert.NotContains(t, callErr.Error(), "upstream.example") assert.NotContains(t, callErr.Error(), "upstream.example")
}) })
t.Run("renders lyrics for an alias on retrieval without upstreamModel", func(t *testing.T) {
lyricsTask := map[string]any{
"task_id": "lyrics-alias", "status": "SUCCESS", "data": map[string]any{"id": "lyrics-1", "title": "Tide", "text": "Sea lyrics"},
}
finalValue := callProtocol(t, "renderFinal", map[string]any{"model": "my-lyrics"}, lyricsTask)
machine := relay.NewPluginResponsesMachine("lyrics-alias", "my-lyrics", 10, relay.DefaultPluginProtocolLimits())
response, finalErr := machine.FinalResponse(finalValue, "SUCCESS")
require.NoError(t, finalErr)
output, ok := response["output"].([]any)
require.True(t, ok)
message, ok := output[0].(map[string]any)
require.True(t, ok)
content, ok := message["content"].([]any)
require.True(t, ok)
require.Len(t, content, 1)
part, ok := content[0].(map[string]any)
require.True(t, ok)
assert.Contains(t, part["text"], "Sea lyrics")
eventsValue := callProtocol(t, "renderEvents", map[string]any{"model": "my-lyrics"}, lyricsTask)
events, decodeErr := relay.DecodePluginProtocolEventResult(eventsValue, relay.DefaultPluginProtocolLimits())
require.NoError(t, decodeErr)
require.Len(t, events.Events, 1)
var text string
require.NoError(t, common.Unmarshal(events.Events[0].Data, &text))
assert.Contains(t, text, "Sea lyrics")
})
} }
...@@ -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.1", version: "1.0.2",
author: { name: "QuantumNous" }, author: { name: "QuantumNous" },
channelTypes: [36], channelTypes: [36],
models: ["suno_music", "suno_lyrics"], models: ["suno_music", "suno_lyrics"],
...@@ -102,14 +102,13 @@ function validateAndNormalize(ctx) { ...@@ -102,14 +102,13 @@ function validateAndNormalize(ctx) {
export function buildSubmitRequest(ctx) { export function buildSubmitRequest(ctx) {
const normalized = validateAndNormalize(ctx); const normalized = validateAndNormalize(ctx);
const incoming = ctx.requestHeaders || {}; const incoming = ctx.requestHeaders || {};
const headers = { Authorization: "Bearer " + ctx.apiKey };
if (typeof incoming["Content-Type"] === "string" && incoming["Content-Type"]) headers["Content-Type"] = incoming["Content-Type"];
if (typeof incoming.Accept === "string" && incoming.Accept) headers.Accept = incoming.Accept;
return { return {
url: ctx.baseUrl + "/suno/submit/" + normalized.action, url: ctx.baseUrl + "/suno/submit/" + normalized.action,
method: "POST", method: "POST",
headers: { headers: headers,
"Content-Type": incoming["Content-Type"] || "",
Accept: incoming.Accept || "",
Authorization: "Bearer " + ctx.apiKey,
},
body: normalized.body, body: normalized.body,
action: normalized.action, action: normalized.action,
}; };
...@@ -228,7 +227,7 @@ function escapedAttribute(value) { ...@@ -228,7 +227,7 @@ function escapedAttribute(value) {
} }
function responseContent(ctx, task) { function responseContent(ctx, task) {
const model = trimmed(ctx && ctx.requestBody && ctx.requestBody.model).toLowerCase(); const declared = trimmed(ctx.upstreamModel || ctx.model).toLowerCase();
const songs = artifactData(task); const songs = artifactData(task);
const lyrics = []; const lyrics = [];
for (const song of songs) { for (const song of songs) {
...@@ -238,7 +237,13 @@ function responseContent(ctx, task) { ...@@ -238,7 +237,13 @@ function responseContent(ctx, task) {
const title = trimmed(song.title); const title = trimmed(song.title);
lyrics.push(title ? title + "\n" + text : text); lyrics.push(title ? title + "\n" + text : text);
} }
if (model === "suno_lyrics") { const lyricsMode =
declared === "suno_lyrics" ||
(declared !== "suno_music" &&
!songs.some(function (song) {
return song && trimmed(song.audio_url);
}));
if (lyricsMode) {
return [{ type: "output_text", text: lyrics.join("\n\n") || "Lyrics generation completed.", annotations: [], logprobs: [] }]; return [{ type: "output_text", text: lyrics.join("\n\n") || "Lyrics generation completed.", annotations: [], logprobs: [] }];
} }
const content = [{ type: "output_text", text: lyrics.join("\n\n") || "Music generation completed.", annotations: [], logprobs: [] }]; const content = [{ type: "output_text", text: lyrics.join("\n\n") || "Music generation completed.", annotations: [], logprobs: [] }];
...@@ -267,21 +272,21 @@ export const protocols = { ...@@ -267,21 +272,21 @@ export const protocols = {
if (!ctx.body || ctx.body.kind !== "json") throw new Error("JSON body required"); if (!ctx.body || ctx.body.kind !== "json") throw new Error("JSON body required");
const req = ctx.body.value; const req = ctx.body.value;
if (!req || typeof req !== "object" || Array.isArray(req)) throw new Error("request body must be an object"); if (!req || typeof req !== "object" || Array.isArray(req)) throw new Error("request body must be an object");
const model = trimmed(req.model); const declared = trimmed(ctx.upstreamModel || ctx.model).toLowerCase();
if (model !== "suno_music" && model !== "suno_lyrics") throw new Error("model is required"); if (declared !== "suno_music" && declared !== "suno_lyrics") throw new Error("model must be suno_music or suno_lyrics");
if (req.input !== undefined && typeof req.input !== "string" && !Array.isArray(req.input)) throw new Error("input must be a string or array"); if (req.input !== undefined && typeof req.input !== "string" && !Array.isArray(req.input)) throw new Error("input must be a string or array");
if (req.metadata !== undefined && (!req.metadata || typeof req.metadata !== "object" || Array.isArray(req.metadata))) if (req.metadata !== undefined && (!req.metadata || typeof req.metadata !== "object" || Array.isArray(req.metadata)))
throw new Error("metadata must be an object"); throw new Error("metadata must be an object");
const input = responsesText(req); const input = responsesText(req);
const requestBody = Object.assign({}, req.metadata || {}); const requestBody = Object.assign({}, req.metadata || {});
if (model === "suno_lyrics") { if (declared === "suno_lyrics") {
if (!trimmed(requestBody.prompt)) requestBody.prompt = input || trimmed(req.prompt); if (!trimmed(requestBody.prompt)) requestBody.prompt = input || trimmed(req.prompt);
if (!trimmed(requestBody.prompt)) throw new Error("input is required"); if (!trimmed(requestBody.prompt)) throw new Error("input is required");
return { kind: "submit", model: model, action: "LYRICS", requestBody: requestBody }; return { kind: "submit", model: ctx.model, action: "LYRICS", requestBody: requestBody };
} }
if (!trimmed(requestBody.gpt_description_prompt)) requestBody.gpt_description_prompt = input || trimmed(req.prompt); if (!trimmed(requestBody.gpt_description_prompt)) requestBody.gpt_description_prompt = input || trimmed(req.prompt);
if (!trimmed(requestBody.gpt_description_prompt) && !trimmed(requestBody.prompt)) throw new Error("input is required"); if (!trimmed(requestBody.gpt_description_prompt) && !trimmed(requestBody.prompt)) throw new Error("input is required");
return { kind: "submit", model: model, action: "MUSIC", requestBody: requestBody }; return { kind: "submit", model: ctx.model, action: "MUSIC", requestBody: requestBody };
}, },
renderEvents: function (ctx, task, previousState) { renderEvents: function (ctx, task, previousState) {
const status = String(task.status || "UNKNOWN").toUpperCase(); const status = String(task.status || "UNKNOWN").toUpperCase();
......
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