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:
@@ -580,12 +580,20 @@ func (s *Server) listModels(w http.ResponseWriter, r *http.Request) {
|
||||
// @Tags playground
|
||||
// @Produce json
|
||||
// @Security BearerAuth
|
||||
// @Param usage_scene query string false "模型可选场景;不传时保持原有行为" Enums(canvas_model_node,desktop)
|
||||
// @Success 200 {object} PlatformModelListResponse
|
||||
// @Failure 400 {object} ErrorEnvelope
|
||||
// @Failure 401 {object} ErrorEnvelope
|
||||
// @Failure 502 {object} ErrorEnvelope
|
||||
// @Failure 500 {object} ErrorEnvelope
|
||||
// @Router /api/v1/models [get]
|
||||
// @Router /api/v1/playground/models [get]
|
||||
func (s *Server) listPlayableModels(w http.ResponseWriter, r *http.Request) {
|
||||
usageScene, err := parseModelUsageSceneQuery(r.URL.Query())
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, err.Error())
|
||||
return
|
||||
}
|
||||
user, _ := auth.UserFromContext(r.Context())
|
||||
models, err := s.store.ListAccessiblePlatformModels(r.Context(), user)
|
||||
if err != nil {
|
||||
@@ -593,6 +601,15 @@ func (s *Server) listPlayableModels(w http.ResponseWriter, r *http.Request) {
|
||||
writeError(w, http.StatusInternalServerError, "list playable models failed")
|
||||
return
|
||||
}
|
||||
if usageScene != "" {
|
||||
availableModelIDs, sceneErr := s.listServerMainUsageSceneModelIDs(r.Context(), usageScene)
|
||||
if sceneErr != nil {
|
||||
s.logger.Error("filter playable models by usage scene failed", "usage_scene", usageScene, "error", sceneErr)
|
||||
writeError(w, http.StatusBadGateway, "model usage scene catalog is unavailable")
|
||||
return
|
||||
}
|
||||
models = filterPlatformModelsByUsageSceneIDs(models, availableModelIDs)
|
||||
}
|
||||
writeJSON(w, http.StatusOK, map[string]any{"items": s.platformModelResponses(r.Context(), models)})
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,186 @@
|
||||
package httpapi
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||
)
|
||||
|
||||
const maxUsageSceneCatalogResponseBytes = 8 << 20
|
||||
|
||||
var supportedModelUsageScenes = map[string]struct{}{
|
||||
"canvas_model_node": {},
|
||||
"desktop": {},
|
||||
}
|
||||
|
||||
func parseModelUsageSceneQuery(query url.Values) (string, error) {
|
||||
values, present := query["usage_scene"]
|
||||
if !present {
|
||||
return "", nil
|
||||
}
|
||||
if len(values) != 1 {
|
||||
return "", errors.New("usage_scene must be provided exactly once")
|
||||
}
|
||||
scene := strings.TrimSpace(values[0])
|
||||
if _, ok := supportedModelUsageScenes[scene]; !ok {
|
||||
return "", errors.New("usage_scene must be canvas_model_node or desktop")
|
||||
}
|
||||
return scene, nil
|
||||
}
|
||||
|
||||
func (s *Server) listServerMainUsageSceneModelIDs(ctx context.Context, usageScene string) (map[string]struct{}, error) {
|
||||
baseURL := strings.TrimRight(strings.TrimSpace(s.cfg.ServerMainBaseURL), "/")
|
||||
if baseURL == "" {
|
||||
return nil, errors.New("SERVER_MAIN_BASE_URL is not configured")
|
||||
}
|
||||
var requestErrors []error
|
||||
for _, endpoint := range serverMainUsageSceneCatalogEndpoints(baseURL, usageScene) {
|
||||
ids, err := s.requestServerMainUsageSceneModelIDs(ctx, endpoint)
|
||||
if err == nil {
|
||||
return ids, nil
|
||||
}
|
||||
requestErrors = append(requestErrors, err)
|
||||
}
|
||||
return nil, fmt.Errorf("server-main model catalog request failed: %w", errors.Join(requestErrors...))
|
||||
}
|
||||
|
||||
func serverMainUsageSceneCatalogEndpoints(baseURL string, usageScene string) []string {
|
||||
path := "/integration-platform/models/strict/all?usage_scene=" + url.QueryEscape(usageScene)
|
||||
if strings.HasSuffix(baseURL, "/api") {
|
||||
return []string{baseURL + path}
|
||||
}
|
||||
return []string{baseURL + "/api" + path, baseURL + path}
|
||||
}
|
||||
|
||||
func (s *Server) requestServerMainUsageSceneModelIDs(ctx context.Context, endpoint string) (map[string]struct{}, error) {
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create server-main model catalog request: %w", err)
|
||||
}
|
||||
request.Header.Set("Accept", "application/json")
|
||||
|
||||
client := http.DefaultClient
|
||||
if s.auth != nil && s.auth.HTTPClient != nil {
|
||||
client = s.auth.HTTPClient
|
||||
}
|
||||
response, err := client.Do(request)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("request server-main model catalog: %w", err)
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices {
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 4<<10))
|
||||
return nil, fmt.Errorf("server-main model catalog returned HTTP %d", response.StatusCode)
|
||||
}
|
||||
|
||||
var payload any
|
||||
decoder := json.NewDecoder(io.LimitReader(response.Body, maxUsageSceneCatalogResponseBytes))
|
||||
if err := decoder.Decode(&payload); err != nil {
|
||||
return nil, fmt.Errorf("decode server-main model catalog: %w", err)
|
||||
}
|
||||
items, ok := usageSceneCatalogItems(payload)
|
||||
if !ok {
|
||||
return nil, errors.New("server-main model catalog payload is malformed")
|
||||
}
|
||||
return usageSceneModelIDs(items), nil
|
||||
}
|
||||
|
||||
func usageSceneCatalogItems(payload any) ([]any, bool) {
|
||||
if items, ok := payload.([]any); ok {
|
||||
return items, true
|
||||
}
|
||||
record, ok := payload.(map[string]any)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
for _, key := range []string{"modelList", "items", "data"} {
|
||||
value, exists := record[key]
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
if items, ok := value.([]any); ok {
|
||||
return items, true
|
||||
}
|
||||
if items, ok := usageSceneCatalogItems(value); ok {
|
||||
return items, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func usageSceneModelIDs(items []any) map[string]struct{} {
|
||||
ids := make(map[string]struct{}, len(items))
|
||||
for _, item := range items {
|
||||
record, ok := item.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
for _, key := range []string{
|
||||
"alias",
|
||||
"invocationName",
|
||||
"invocation_name",
|
||||
"modelAlias",
|
||||
"model_alias",
|
||||
"name",
|
||||
"modelName",
|
||||
"model_name",
|
||||
"providerModelName",
|
||||
"provider_model_name",
|
||||
} {
|
||||
addUsageSceneModelID(ids, record[key])
|
||||
}
|
||||
for _, key := range []string{"legacyAliases", "legacy_aliases"} {
|
||||
if aliases, ok := record[key].([]any); ok {
|
||||
for _, alias := range aliases {
|
||||
addUsageSceneModelID(ids, alias)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func filterPlatformModelsByUsageSceneIDs(models []store.PlatformModel, availableModelIDs map[string]struct{}) []store.PlatformModel {
|
||||
filtered := make([]store.PlatformModel, 0, len(models))
|
||||
for _, model := range models {
|
||||
if platformModelMatchesUsageSceneIDs(model, availableModelIDs) {
|
||||
filtered = append(filtered, model)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func platformModelMatchesUsageSceneIDs(model store.PlatformModel, availableModelIDs map[string]struct{}) bool {
|
||||
for _, value := range []string{model.ModelName, model.ProviderModelName, model.ModelAlias} {
|
||||
if _, ok := availableModelIDs[normalizeUsageSceneModelID(value)]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
for _, alias := range model.LegacyAliases {
|
||||
if _, ok := availableModelIDs[normalizeUsageSceneModelID(alias)]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func addUsageSceneModelID(ids map[string]struct{}, value any) {
|
||||
text, ok := value.(string)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if normalized := normalizeUsageSceneModelID(text); normalized != "" {
|
||||
ids[normalized] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeUsageSceneModelID(value string) string {
|
||||
return strings.ToLower(strings.TrimSpace(value))
|
||||
}
|
||||
@@ -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