Files
easyai-ai-gateway/apps/api/internal/httpapi/model_usage_scene_test.go
T
chengcheng 31c32690b2 feat(models): 支持按使用场景筛选模型
Gateway 模型接口校验 usage_scene,并与 server-main 严格场景目录取交集;显式场景查询在上游不可用时关闭失败,避免返回不应暴露的模型。\n\n验证:env -u AI_GATEWAY_TEST_DATABASE_URL go test ./... -count=1
2026-07-27 13:15:23 +08:00

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)
}
}