Gateway 模型接口校验 usage_scene,并与 server-main 严格场景目录取交集;显式场景查询在上游不可用时关闭失败,避免返回不应暴露的模型。\n\n验证:env -u AI_GATEWAY_TEST_DATABASE_URL go test ./... -count=1
161 lines
5.3 KiB
Go
161 lines
5.3 KiB
Go
package httpapi
|
|
|
|
import (
|
|
"context"
|
|
"io"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"net/url"
|
|
"testing"
|
|
|
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
|
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
|
)
|
|
|
|
func TestParseModelUsageSceneQuery(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
for _, testCase := range []struct {
|
|
name string
|
|
query url.Values
|
|
want string
|
|
wantErr bool
|
|
}{
|
|
{name: "missing", query: url.Values{}, want: ""},
|
|
{name: "desktop", query: url.Values{"usage_scene": {"desktop"}}, want: "desktop"},
|
|
{name: "canvas", query: url.Values{"usage_scene": {"canvas_model_node"}}, want: "canvas_model_node"},
|
|
{name: "empty", query: url.Values{"usage_scene": {""}}, wantErr: true},
|
|
{name: "unknown", query: url.Values{"usage_scene": {"mobile"}}, wantErr: true},
|
|
{name: "duplicate", query: url.Values{"usage_scene": {"desktop", "canvas_model_node"}}, wantErr: true},
|
|
} {
|
|
testCase := testCase
|
|
t.Run(testCase.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
got, err := parseModelUsageSceneQuery(testCase.query)
|
|
if testCase.wantErr {
|
|
if err == nil {
|
|
t.Fatal("expected usage_scene validation error")
|
|
}
|
|
return
|
|
}
|
|
if err != nil {
|
|
t.Fatalf("parse usage_scene: %v", err)
|
|
}
|
|
if got != testCase.want {
|
|
t.Fatalf("usage_scene = %q, want %q", got, testCase.want)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestListServerMainUsageSceneModelIDs(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
serverMain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path != "/api/integration-platform/models/strict/all" {
|
|
t.Errorf("path = %q", r.URL.Path)
|
|
}
|
|
if got := r.URL.Query().Get("usage_scene"); got != "desktop" {
|
|
t.Errorf("usage_scene = %q", got)
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = io.WriteString(w, `{"data":[{"alias":"Qwen3.7-Plus"},{"invocationName":"GPT-5"},{"legacyAliases":["gpt-latest"]}]}`)
|
|
}))
|
|
defer serverMain.Close()
|
|
|
|
authenticator := auth.New("secret", serverMain.URL, "")
|
|
gateway := &Server{
|
|
cfg: config.Config{ServerMainBaseURL: serverMain.URL},
|
|
auth: authenticator,
|
|
logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
|
|
}
|
|
ids, err := gateway.listServerMainUsageSceneModelIDs(context.Background(), "desktop")
|
|
if err != nil {
|
|
t.Fatalf("list server-main usage scene models: %v", err)
|
|
}
|
|
for _, id := range []string{"qwen3.7-plus", "gpt-5", "gpt-latest"} {
|
|
if _, ok := ids[id]; !ok {
|
|
t.Errorf("missing normalized model ID %q in %#v", id, ids)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestFilterPlatformModelsByUsageSceneIDs(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
models := []store.PlatformModel{
|
|
{ID: "allowed-by-invocation", ModelName: "Qwen3.7-Plus", ProviderModelName: "qwen3.7-plus-2026-07-01"},
|
|
{ID: "allowed-by-provider-name", ModelName: "GPT-5", ProviderModelName: "gpt-5-2026-07-01"},
|
|
{ID: "allowed-by-legacy-alias", ModelName: "renamed-model", LegacyAliases: store.StringList{"old-model"}},
|
|
{ID: "desktop-restricted", ModelName: "Qwen3.6-Plus", ProviderModelName: "qwen3.6-plus-2026-04-02"},
|
|
}
|
|
allowed := map[string]struct{}{
|
|
"qwen3.7-plus": {},
|
|
"gpt-5-2026-07-01": {},
|
|
"old-model": {},
|
|
}
|
|
|
|
filtered := filterPlatformModelsByUsageSceneIDs(models, allowed)
|
|
if len(filtered) != 3 {
|
|
t.Fatalf("filtered model count = %d, want 3: %#v", len(filtered), filtered)
|
|
}
|
|
for _, model := range filtered {
|
|
if model.ID == "desktop-restricted" {
|
|
t.Fatalf("desktop-restricted model remained in filtered catalog: %#v", filtered)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestListServerMainUsageSceneModelIDsRejectsMalformedPayload(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
serverMain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = io.WriteString(w, `{"success":true}`)
|
|
}))
|
|
defer serverMain.Close()
|
|
|
|
gateway := &Server{
|
|
cfg: config.Config{ServerMainBaseURL: serverMain.URL},
|
|
auth: auth.New("secret", serverMain.URL, ""),
|
|
}
|
|
if _, err := gateway.listServerMainUsageSceneModelIDs(context.Background(), "desktop"); err == nil {
|
|
t.Fatal("expected malformed upstream payload error")
|
|
}
|
|
}
|
|
|
|
func TestListServerMainUsageSceneModelIDsFallsBackToDirectServerMainRoute(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
requestedPaths := make([]string, 0, 2)
|
|
serverMain := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
requestedPaths = append(requestedPaths, r.URL.Path)
|
|
if r.URL.Path == "/api/integration-platform/models/strict/all" {
|
|
http.NotFound(w, r)
|
|
return
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = io.WriteString(w, `[{"alias":"Qwen3.7-Plus"}]`)
|
|
}))
|
|
defer serverMain.Close()
|
|
|
|
gateway := &Server{
|
|
cfg: config.Config{ServerMainBaseURL: serverMain.URL},
|
|
auth: auth.New("secret", serverMain.URL, ""),
|
|
}
|
|
ids, err := gateway.listServerMainUsageSceneModelIDs(context.Background(), "desktop")
|
|
if err != nil {
|
|
t.Fatalf("list direct server-main usage scene models: %v", err)
|
|
}
|
|
if _, ok := ids["qwen3.7-plus"]; !ok {
|
|
t.Fatalf("missing direct server-main model ID: %#v", ids)
|
|
}
|
|
if len(requestedPaths) != 2 ||
|
|
requestedPaths[0] != "/api/integration-platform/models/strict/all" ||
|
|
requestedPaths[1] != "/integration-platform/models/strict/all" {
|
|
t.Fatalf("unexpected server-main route probes: %#v", requestedPaths)
|
|
}
|
|
}
|