From 31c32690b201265024ab6083168451f82b0bd66b Mon Sep 17 00:00:00 2001 From: chengcheng Date: Mon, 27 Jul 2026 13:15:23 +0800 Subject: [PATCH] =?UTF-8?q?feat(models):=20=E6=94=AF=E6=8C=81=E6=8C=89?= =?UTF-8?q?=E4=BD=BF=E7=94=A8=E5=9C=BA=E6=99=AF=E7=AD=9B=E9=80=89=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Gateway 模型接口校验 usage_scene,并与 server-main 严格场景目录取交集;显式场景查询在上游不可用时关闭失败,避免返回不应暴露的模型。\n\n验证:env -u AI_GATEWAY_TEST_DATABASE_URL go test ./... -count=1 --- apps/api/docs/swagger.json | 48 +++++ apps/api/docs/swagger.yaml | 32 +++ apps/api/internal/httpapi/handlers.go | 17 ++ .../api/internal/httpapi/model_usage_scene.go | 186 ++++++++++++++++++ .../httpapi/model_usage_scene_test.go | 160 +++++++++++++++ 5 files changed, 443 insertions(+) create mode 100644 apps/api/internal/httpapi/model_usage_scene.go create mode 100644 apps/api/internal/httpapi/model_usage_scene_test.go diff --git a/apps/api/docs/swagger.json b/apps/api/docs/swagger.json index 2a466eb..e4f5bc9 100644 --- a/apps/api/docs/swagger.json +++ b/apps/api/docs/swagger.json @@ -6860,6 +6860,18 @@ "playground" ], "summary": "列出可调用模型", + "parameters": [ + { + "enum": [ + "canvas_model_node", + "desktop" + ], + "type": "string", + "description": "模型可选场景;不传时保持原有行为", + "name": "usage_scene", + "in": "query" + } + ], "responses": { "200": { "description": "OK", @@ -6867,6 +6879,12 @@ "$ref": "#/definitions/httpapi.PlatformModelListResponse" } }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/httpapi.ErrorEnvelope" + } + }, "401": { "description": "Unauthorized", "schema": { @@ -6878,6 +6896,12 @@ "schema": { "$ref": "#/definitions/httpapi.ErrorEnvelope" } + }, + "502": { + "description": "Bad Gateway", + "schema": { + "$ref": "#/definitions/httpapi.ErrorEnvelope" + } } } } @@ -7220,6 +7244,18 @@ "playground" ], "summary": "列出可调用模型", + "parameters": [ + { + "enum": [ + "canvas_model_node", + "desktop" + ], + "type": "string", + "description": "模型可选场景;不传时保持原有行为", + "name": "usage_scene", + "in": "query" + } + ], "responses": { "200": { "description": "OK", @@ -7227,6 +7263,12 @@ "$ref": "#/definitions/httpapi.PlatformModelListResponse" } }, + "400": { + "description": "Bad Request", + "schema": { + "$ref": "#/definitions/httpapi.ErrorEnvelope" + } + }, "401": { "description": "Unauthorized", "schema": { @@ -7238,6 +7280,12 @@ "schema": { "$ref": "#/definitions/httpapi.ErrorEnvelope" } + }, + "502": { + "description": "Bad Gateway", + "schema": { + "$ref": "#/definitions/httpapi.ErrorEnvelope" + } } } } diff --git a/apps/api/docs/swagger.yaml b/apps/api/docs/swagger.yaml index 3b68a48..514e8cb 100644 --- a/apps/api/docs/swagger.yaml +++ b/apps/api/docs/swagger.yaml @@ -8381,6 +8381,14 @@ paths: /api/v1/models: get: description: 按当前用户权限返回可用于 Playground 或 API 调用的模型列表。 + parameters: + - description: 模型可选场景;不传时保持原有行为 + enum: + - canvas_model_node + - desktop + in: query + name: usage_scene + type: string produces: - application/json responses: @@ -8388,6 +8396,10 @@ paths: description: OK schema: $ref: '#/definitions/httpapi.PlatformModelListResponse' + "400": + description: Bad Request + schema: + $ref: '#/definitions/httpapi.ErrorEnvelope' "401": description: Unauthorized schema: @@ -8396,6 +8408,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/httpapi.ErrorEnvelope' + "502": + description: Bad Gateway + schema: + $ref: '#/definitions/httpapi.ErrorEnvelope' security: - BearerAuth: [] summary: 列出可调用模型 @@ -8615,6 +8631,14 @@ paths: /api/v1/playground/models: get: description: 按当前用户权限返回可用于 Playground 或 API 调用的模型列表。 + parameters: + - description: 模型可选场景;不传时保持原有行为 + enum: + - canvas_model_node + - desktop + in: query + name: usage_scene + type: string produces: - application/json responses: @@ -8622,6 +8646,10 @@ paths: description: OK schema: $ref: '#/definitions/httpapi.PlatformModelListResponse' + "400": + description: Bad Request + schema: + $ref: '#/definitions/httpapi.ErrorEnvelope' "401": description: Unauthorized schema: @@ -8630,6 +8658,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/httpapi.ErrorEnvelope' + "502": + description: Bad Gateway + schema: + $ref: '#/definitions/httpapi.ErrorEnvelope' security: - BearerAuth: [] summary: 列出可调用模型 diff --git a/apps/api/internal/httpapi/handlers.go b/apps/api/internal/httpapi/handlers.go index 1be912e..2c70dd1 100644 --- a/apps/api/internal/httpapi/handlers.go +++ b/apps/api/internal/httpapi/handlers.go @@ -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)}) } diff --git a/apps/api/internal/httpapi/model_usage_scene.go b/apps/api/internal/httpapi/model_usage_scene.go new file mode 100644 index 0000000..07bbdbe --- /dev/null +++ b/apps/api/internal/httpapi/model_usage_scene.go @@ -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)) +} diff --git a/apps/api/internal/httpapi/model_usage_scene_test.go b/apps/api/internal/httpapi/model_usage_scene_test.go new file mode 100644 index 0000000..b3ba35b --- /dev/null +++ b/apps/api/internal/httpapi/model_usage_scene_test.go @@ -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) + } +}