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:
chengcheng
2026-07-27 13:15:23 +08:00
parent de10875439
commit 31c32690b2
5 changed files with 443 additions and 0 deletions
+17
View File
@@ -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)
}
}