mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-01 19:41:57 +00:00
373 lines
12 KiB
Go
373 lines
12 KiB
Go
package jsplugin
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"errors"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
"unicode/utf8"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
func TestEngineCallsESMExportWithInjectedUtils(t *testing.T) {
|
|
t.Parallel()
|
|
logs := make([]string, 0, 1)
|
|
engine, err := Compile(`
|
|
export function sign(ctx) {
|
|
console.log("called", ctx.name);
|
|
return {
|
|
now: utils.unixNow(),
|
|
digest: utils.hmacSHA256(ctx.message, ctx.secret),
|
|
encoded: utils.base64(ctx.message),
|
|
};
|
|
}
|
|
export const meta = { apiVersion: 1, key: "mock" };
|
|
`, Options{
|
|
Key: "mock", Version: "1.0.0",
|
|
Now: func() time.Time { return time.Unix(1234, 0) },
|
|
Log: func(message string) { logs = append(logs, message) },
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
result, err := engine.Call(context.Background(), "sign", map[string]any{
|
|
"name": "fixture", "message": "hello", "secret": "secret",
|
|
})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, map[string]any{
|
|
"now": int64(1234), "digest": "88aab3ede8d3adf94d26ab90d3bafd4a2083070c3bcce9c014ee04a443847c0b", "encoded": "aGVsbG8=",
|
|
}, result)
|
|
assert.Equal(t, []string{"[plugin:mock@1.0.0] called fixture"}, logs)
|
|
|
|
meta, err := engine.Export(context.Background(), "meta")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, map[string]any{"apiVersion": int64(1), "key": "mock"}, meta)
|
|
}
|
|
|
|
func TestEngineConsoleLogUsesDebugLoggerAndRequestContext(t *testing.T) {
|
|
previousDebug := common.DebugEnabled
|
|
common.DebugEnabled = false
|
|
t.Cleanup(func() { common.DebugEnabled = previousDebug })
|
|
|
|
var output bytes.Buffer
|
|
common.LogWriterMu.Lock()
|
|
previousWriter := gin.DefaultErrorWriter
|
|
gin.DefaultErrorWriter = &output
|
|
common.LogWriterMu.Unlock()
|
|
t.Cleanup(func() {
|
|
common.LogWriterMu.Lock()
|
|
gin.DefaultErrorWriter = previousWriter
|
|
common.LogWriterMu.Unlock()
|
|
})
|
|
|
|
plugin, err := CompilePlugin(`
|
|
export const meta = {
|
|
apiVersion: 1,
|
|
key: "console-debug",
|
|
name: "Console debug",
|
|
version: "1.2.3",
|
|
author: {name: "Test"},
|
|
models: ["debug-model"],
|
|
fetchMode: "per_task",
|
|
};
|
|
export function run(label) {
|
|
console.log("checkpoint", label);
|
|
return true;
|
|
}
|
|
export function buildSubmitRequest() { return {url: "https://example.com"}; }
|
|
export function parseSubmitResponse() { return {taskId: "one"}; }
|
|
export function buildQueryRequest() { return {url: "https://example.com"}; }
|
|
export function parseTaskResult() { return {status: "SUCCESS"}; }
|
|
`, Options{})
|
|
require.NoError(t, err)
|
|
engine := plugin.Engine
|
|
|
|
disabledContext := context.WithValue(context.Background(), common.RequestIdKey, "plugin-console-disabled")
|
|
_, err = engine.Call(disabledContext, "run", "disabled")
|
|
require.NoError(t, err)
|
|
assert.Empty(t, output.String())
|
|
|
|
common.DebugEnabled = true
|
|
contextA := context.WithValue(context.Background(), common.RequestIdKey, "plugin-console-request-a")
|
|
_, err = engine.Call(contextA, "run", "context-a")
|
|
require.NoError(t, err)
|
|
logA := output.String()
|
|
output.Reset()
|
|
|
|
contextB := context.WithValue(context.Background(), common.RequestIdKey, "plugin-console-request-b")
|
|
_, err = engine.Call(contextB, "run", "context-b")
|
|
require.NoError(t, err)
|
|
logB := output.String()
|
|
output.Reset()
|
|
|
|
_, err = engine.Call(context.Background(), "run", "background")
|
|
require.NoError(t, err)
|
|
logBackground := output.String()
|
|
|
|
assert.Contains(t, logA, "plugin-console-request-a")
|
|
assert.NotContains(t, logA, "plugin-console-request-b")
|
|
assert.Contains(t, logA, "task_plugin subsystem=runtime event=console")
|
|
assert.Contains(t, logA, "[plugin:console-debug@1.2.3] checkpoint context-a")
|
|
assert.NotContains(t, logA, "disabled")
|
|
|
|
assert.Contains(t, logB, "plugin-console-request-b")
|
|
assert.NotContains(t, logB, "plugin-console-request-a")
|
|
assert.Contains(t, logB, "[plugin:console-debug@1.2.3] checkpoint context-b")
|
|
|
|
assert.Contains(t, logBackground, "| SYSTEM |")
|
|
assert.NotContains(t, logBackground, "plugin-console-request-a")
|
|
assert.NotContains(t, logBackground, "plugin-console-request-b")
|
|
assert.Contains(t, logBackground, "[plugin:console-debug@1.2.3] checkpoint background")
|
|
}
|
|
|
|
func TestCompileRejectsAsynchronousAndImportedPlugins(t *testing.T) {
|
|
t.Parallel()
|
|
for name, source := range map[string]string{
|
|
"async": `export async function run() {}`,
|
|
"static import": `import value from "dependency"; export function run() { return value; }`,
|
|
"dynamic import": `export function run() { return import("dependency"); }`,
|
|
"top-level await": `const value = await work(); export function run() { return value; }`,
|
|
} {
|
|
t.Run(name, func(t *testing.T) {
|
|
_, err := Compile(source, Options{Key: "invalid"})
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "unsupported plugin syntax")
|
|
})
|
|
}
|
|
|
|
_, err := Compile(`export function run() { return "import async await"; }`, Options{Key: "valid"})
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
func TestCompileIgnoresSourceMapDirectives(t *testing.T) {
|
|
t.Parallel()
|
|
// A sourceMappingURL comment must stay inert. Sobek's default loader
|
|
// os.ReadFiles the referenced server path during Compile and turns any
|
|
// load failure into a compile error, so an unresolvable path compiling
|
|
// cleanly proves the loader is disabled.
|
|
engine, err := Compile("export function run() { return 1; }\n//# sourceMappingURL=/nonexistent/leak-probe.map\n", Options{Key: "sourcemap"})
|
|
require.NoError(t, err)
|
|
|
|
value, err := engine.Call(context.Background(), "run")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, int64(1), value)
|
|
}
|
|
|
|
func TestEngineInterruptsLongRunningHook(t *testing.T) {
|
|
t.Parallel()
|
|
engine, err := Compile(`export function run() { while (true) {} }`, Options{
|
|
Key: "loop", Version: "1", Timeout: 20 * time.Millisecond,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
_, err = engine.Call(context.Background(), "run")
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "timed out")
|
|
var hookErr *HookError
|
|
assert.False(t, errors.As(err, &hookErr), "timeouts must not be HookError")
|
|
}
|
|
|
|
func TestEngineHookErrorExtractsSanitizedJSMessage(t *testing.T) {
|
|
t.Parallel()
|
|
tests := []struct {
|
|
name string
|
|
source string
|
|
wantMessage string
|
|
wantLen int
|
|
}{
|
|
{
|
|
name: "Error object",
|
|
source: `export function run() { throw new Error("model is required"); }`,
|
|
wantMessage: "model is required",
|
|
},
|
|
{
|
|
name: "raw string throw",
|
|
source: `export function run() { throw "raw string"; }`,
|
|
wantMessage: "raw string",
|
|
},
|
|
{
|
|
name: "truncates to 512 runes",
|
|
source: `export function run() { throw new Error("x".repeat(2000)); }`,
|
|
wantLen: 512,
|
|
wantMessage: strings.Repeat("x", 512),
|
|
},
|
|
{
|
|
name: "scrubs control characters",
|
|
source: "export function run() { throw new Error(\"line1\\nline2\\x1b[31mred\"); }",
|
|
wantMessage: "line1 line2 [31mred",
|
|
},
|
|
{
|
|
name: "throwing message getter falls back without crashing",
|
|
source: `export function run() { throw {get message() { throw {get message() { return "deep"; }}; }}; }`,
|
|
wantMessage: "plugin hook failed",
|
|
},
|
|
}
|
|
for _, testCase := range tests {
|
|
t.Run(testCase.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
engine, err := Compile(testCase.source, Options{Key: "diag", Version: "1.0.0"})
|
|
require.NoError(t, err)
|
|
|
|
_, err = engine.Call(context.Background(), "run")
|
|
require.Error(t, err)
|
|
|
|
var hookErr *HookError
|
|
require.True(t, errors.As(err, &hookErr))
|
|
assert.Equal(t, "run", hookErr.Hook)
|
|
assert.Equal(t, testCase.wantMessage, hookErr.Message)
|
|
if testCase.wantLen > 0 {
|
|
assert.Equal(t, testCase.wantLen, utf8.RuneCountInString(hookErr.Message))
|
|
}
|
|
assert.Contains(t, hookErr.Error(), "plugin diag@1.0.0")
|
|
assert.NotContains(t, hookErr.Message, "Error:")
|
|
assert.NotContains(t, hookErr.Message, "plugin diag@")
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestEngineReportsProtocolAdmissionTimeoutSeparately(t *testing.T) {
|
|
engine, err := Compile(`
|
|
export const protocols = {
|
|
responses: {renderEvents: function() { return {events: [], done: false}; }},
|
|
};
|
|
`, Options{Key: "admission", Version: "1", Concurrency: 1})
|
|
require.NoError(t, err)
|
|
|
|
engine.semaphore <- struct{}{}
|
|
_, err = engine.CallPathWithAdmissionTimeout(
|
|
context.Background(),
|
|
time.Nanosecond,
|
|
"protocols",
|
|
[]string{"responses", "renderEvents"},
|
|
)
|
|
<-engine.semaphore
|
|
|
|
require.ErrorIs(t, err, ErrCallAdmissionTimeout)
|
|
result, err := engine.CallPathWithAdmissionTimeout(
|
|
context.Background(),
|
|
time.Second,
|
|
"protocols",
|
|
[]string{"responses", "renderEvents"},
|
|
)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, map[string]any{"events": []any{}, "done": false}, result)
|
|
}
|
|
|
|
func TestEngineExportInterruptsLongRunningGetter(t *testing.T) {
|
|
t.Parallel()
|
|
engine, err := Compile(`
|
|
export const meta = {
|
|
apiVersion: 1,
|
|
get name() { while (true) {} },
|
|
};
|
|
`, Options{Key: "meta-loop", Version: "1.0.0", Timeout: 20 * time.Millisecond})
|
|
require.NoError(t, err)
|
|
|
|
_, err = engine.Export(context.Background(), "meta")
|
|
require.ErrorContains(t, err, "export meta interrupted")
|
|
}
|
|
|
|
func TestEngineExportReturnsThrownGetterError(t *testing.T) {
|
|
t.Parallel()
|
|
engine, err := Compile(`
|
|
export const meta = {
|
|
apiVersion: 1,
|
|
get name() { throw new Error("getter failed"); },
|
|
};
|
|
`, Options{Key: "meta-throw", Version: "1.0.0"})
|
|
require.NoError(t, err)
|
|
|
|
_, err = engine.Export(context.Background(), "meta")
|
|
require.ErrorContains(t, err, "export meta failed")
|
|
assert.Contains(t, err.Error(), "getter failed")
|
|
}
|
|
|
|
func TestEngineNestedHooksRequireOwnProperties(t *testing.T) {
|
|
t.Parallel()
|
|
engine, err := Compile(`
|
|
const inheritedRenderers = {
|
|
inherited: function(value) { return value; },
|
|
constructor: function(value) { return value; },
|
|
toString: function(value) { return value; },
|
|
["__proto__"]: function(value) { return value; },
|
|
};
|
|
export const renderers = Object.create(inheritedRenderers);
|
|
renderers.own = function(value) { return {id: value.id}; };
|
|
|
|
const inheritedProtocol = {
|
|
renderFinal: function(value) { return value; },
|
|
};
|
|
export const protocols = {
|
|
responses: Object.create(inheritedProtocol),
|
|
};
|
|
`, Options{Key: "own-hooks", Version: "1.0.0"})
|
|
require.NoError(t, err)
|
|
|
|
found, err := engine.HasCallablePath(context.Background(), "renderers", "own")
|
|
require.NoError(t, err)
|
|
assert.True(t, found)
|
|
result, err := engine.CallMember(context.Background(), "renderers", "own", map[string]any{"id": "task-1"})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, map[string]any{"id": "task-1"}, result)
|
|
|
|
for _, member := range []string{"inherited", "constructor", "toString", "__proto__"} {
|
|
t.Run(member, func(t *testing.T) {
|
|
found, err := engine.HasCallablePath(context.Background(), "renderers", member)
|
|
require.NoError(t, err)
|
|
assert.False(t, found)
|
|
|
|
_, err = engine.CallMember(context.Background(), "renderers", member)
|
|
require.ErrorContains(t, err, "not found")
|
|
})
|
|
}
|
|
|
|
found, err = engine.HasCallablePath(context.Background(), "protocols", "responses", "renderFinal")
|
|
require.NoError(t, err)
|
|
assert.False(t, found)
|
|
_, err = engine.CallPath(context.Background(), "protocols", []string{"responses", "renderFinal"})
|
|
require.ErrorContains(t, err, "not found")
|
|
}
|
|
|
|
func TestCompileInterruptsLongRunningInitialization(t *testing.T) {
|
|
t.Parallel()
|
|
_, err := Compile(`while (true) {}; export function run() {}`, Options{
|
|
Key: "loop", Version: "1", Timeout: 20 * time.Millisecond,
|
|
})
|
|
require.Error(t, err)
|
|
assert.Contains(t, err.Error(), "initialization timed out")
|
|
}
|
|
|
|
func TestValidateRequestURL(t *testing.T) {
|
|
t.Parallel()
|
|
tests := []struct {
|
|
name string
|
|
requestURL string
|
|
baseURL string
|
|
allowedHosts []string
|
|
wantError string
|
|
}{
|
|
{name: "same host", requestURL: "https://api.example.com/v1/task", baseURL: "https://api.example.com/v1"},
|
|
{name: "default port", requestURL: "https://api.example.com:443/v1/task", baseURL: "https://api.example.com"},
|
|
{name: "approved host", requestURL: "https://upload.example.com/task", baseURL: "https://api.example.com", allowedHosts: []string{"upload.example.com"}},
|
|
{name: "subdomain is not implicit", requestURL: "https://evil.api.example.com/task", baseURL: "https://api.example.com", wantError: "not allowed"},
|
|
{name: "userinfo trick", requestURL: "https://api.example.com@evil.example/task", baseURL: "https://api.example.com", wantError: "not allowed"},
|
|
{name: "relative URL", requestURL: "/v1/task", baseURL: "https://api.example.com", wantError: "absolute"},
|
|
}
|
|
for _, test := range tests {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
err := ValidateRequestURL(test.requestURL, test.baseURL, test.allowedHosts)
|
|
if test.wantError == "" {
|
|
require.NoError(t, err)
|
|
return
|
|
}
|
|
require.Error(t, err)
|
|
assert.True(t, strings.Contains(err.Error(), test.wantError), err.Error())
|
|
})
|
|
}
|
|
}
|