Files
easyai-ai-gateway/apps/api/internal/httpapi/model_usage_scene.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

187 lines
5.2 KiB
Go

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