feat(models): 支持按使用场景筛选模型
Gateway 模型接口校验 usage_scene,并与 server-main 严格场景目录取交集;显式场景查询在上游不可用时关闭失败,避免返回不应暴露的模型。\n\n验证:env -u AI_GATEWAY_TEST_DATABASE_URL go test ./... -count=1
This commit is contained in:
@@ -0,0 +1,160 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user