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 (
"io/fs"
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/pkg/jsplugin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
var expectedKeys = []string{"alibaba", "doubao", "google", "hailuo", "jimeng", "kling", "sora", "sunoapi", "vertex-ai", "vidu"}
func TestBuiltInVendorPluginsDeclareNativeRoutesAndLegacyChannelTypes(t *testing.T) {
generation := jsplugin.DefaultRegistry.Generation()
require.NotNil(t, generation)
......@@ -63,7 +66,6 @@ func TestBuiltInVendorPluginsDeclareNativeRoutesAndLegacyChannelTypes(t *testing
}
func TestBuiltInTaskPluginResponsesAndUsageContracts(t *testing.T) {
expectedKeys := []string{"alibaba", "doubao", "google", "hailuo", "jimeng", "kling", "sora", "sunoapi", "vertex-ai", "vidu"}
generation := jsplugin.DefaultRegistry.Generation()
require.NotNil(t, generation)
......@@ -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) {
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) {
value, callErr := plugin.Engine.Call(t.Context(), "extractUsage", map[string]any{
"model": "suno_music", "action": "MUSIC", "usagePurpose": "facts", "requestBody": map[string]any{},
......@@ -92,8 +116,8 @@ func TestSunoResponsesProtocol(t *testing.T) {
const firstAudioKey = "audio-5fc0a0cd3367274b4b6de056fc754263f8726a704bb4814ffeb88495f22dad35"
const secondAudioKey = "audio-a9c04b840373f4ef4e8d80140b745c6f647819fa375bc34368cdccced7e2b455"
protocolContext := map[string]any{
"requestBody": map[string]any{"model": "suno_music"},
"stream": true,
"model": "suno_music",
"stream": true,
"artifacts": map[string]any{
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"},
......@@ -167,7 +191,7 @@ func TestSunoResponsesProtocol(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"},
})
machine := relay.NewPluginResponsesMachine("lyrics-public", "suno_lyrics", 10, relay.DefaultPluginProtocolLimits())
......@@ -187,9 +211,37 @@ func TestSunoResponsesProtocol(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{
"requestBody": map[string]any{"model": "suno_music"},
"model": "suno_music",
}, successTask)
require.ErrorContains(t, callErr, "audio artifact is unavailable")
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 = {
en: "SunoAPI project music and lyrics generation",
zh: "SunoAPI 项目 音乐与歌词生成",
},
version: "1.0.1",
version: "1.0.2",
author: { name: "QuantumNous" },
channelTypes: [36],
models: ["suno_music", "suno_lyrics"],
......@@ -102,14 +102,13 @@ function validateAndNormalize(ctx) {
export function buildSubmitRequest(ctx) {
const normalized = validateAndNormalize(ctx);
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 {
url: ctx.baseUrl + "/suno/submit/" + normalized.action,
method: "POST",
headers: {
"Content-Type": incoming["Content-Type"] || "",
Accept: incoming.Accept || "",
Authorization: "Bearer " + ctx.apiKey,
},
headers: headers,
body: normalized.body,
action: normalized.action,
};
......@@ -228,7 +227,7 @@ function escapedAttribute(value) {
}
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 lyrics = [];
for (const song of songs) {
......@@ -238,7 +237,13 @@ function responseContent(ctx, task) {
const title = trimmed(song.title);
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: [] }];
}
const content = [{ type: "output_text", text: lyrics.join("\n\n") || "Music generation completed.", annotations: [], logprobs: [] }];
......@@ -267,21 +272,21 @@ export const protocols = {
if (!ctx.body || ctx.body.kind !== "json") throw new Error("JSON body required");
const req = ctx.body.value;
if (!req || typeof req !== "object" || Array.isArray(req)) throw new Error("request body must be an object");
const model = trimmed(req.model);
if (model !== "suno_music" && model !== "suno_lyrics") throw new Error("model is required");
const declared = trimmed(ctx.upstreamModel || ctx.model).toLowerCase();
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.metadata !== undefined && (!req.metadata || typeof req.metadata !== "object" || Array.isArray(req.metadata)))
throw new Error("metadata must be an object");
const input = responsesText(req);
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)) 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) && !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) {
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