Add composite group routing

This commit is contained in:
Heatherm Huang
2026-07-23 09:19:24 +08:00
parent fa2da0409e
commit ebc1028771
30 changed files with 640 additions and 67 deletions
+1
View File
@@ -23,6 +23,7 @@ const (
PlatformGemini = "gemini"
PlatformAntigravity = "antigravity"
PlatformGrok = "grok"
PlatformComposite = "composite"
)
// Account type constants
@@ -87,7 +87,7 @@ func NewGroupHandler(adminService service.AdminService, dashboardService *servic
type CreateGroupRequest struct {
Name string `json:"name" binding:"required"`
Description string `json:"description"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok composite"`
RateMultiplier float64 `json:"rate_multiplier"`
IsExclusive bool `json:"is_exclusive"`
SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"`
@@ -144,7 +144,7 @@ type CreateGroupRequest struct {
type UpdateGroupRequest struct {
Name string `json:"name"`
Description *string `json:"description"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok composite"`
RateMultiplier *float64 `json:"rate_multiplier"`
IsExclusive *bool `json:"is_exclusive"`
Status string `json:"status" binding:"omitempty,oneof=active inactive"`
@@ -0,0 +1,48 @@
package handler
import (
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
func ensureCompositeTargetPlatform(c *gin.Context, apiKey *service.APIKey, model string) {
if c == nil || c.Request == nil || apiKey == nil || apiKey.Group == nil || apiKey.Group.Platform != service.PlatformComposite {
return
}
if _, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok {
return
}
if platform, ok := service.DetectModelPlatform(model); ok {
c.Request = c.Request.WithContext(service.WithResolvedTargetPlatform(c.Request.Context(), platform))
}
}
func compositeTargetPlatformAllowed(c *gin.Context, apiKey *service.APIKey, model string, allowed ...string) bool {
if c == nil || c.Request == nil || apiKey == nil || apiKey.Group == nil || apiKey.Group.Platform != service.PlatformComposite {
return true
}
ensureCompositeTargetPlatform(c, apiKey, model)
platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context())
if !ok {
return false
}
for _, allowedPlatform := range allowed {
if platform == allowedPlatform {
return true
}
}
return false
}
func effectiveAPIKeyPlatform(c *gin.Context, apiKey *service.APIKey) string {
if c != nil && c.Request != nil {
if platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok {
return platform
}
}
if apiKey == nil || apiKey.Group == nil {
return ""
}
return apiKey.Group.Platform
}
@@ -0,0 +1,42 @@
package handler
import (
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestCompositeTargetPlatformAllowedResolvesKnownAllowedModel(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest("POST", "/v1/embeddings", nil)
apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite}}
require.True(t, compositeTargetPlatformAllowed(c, apiKey, "text-embedding-3-large", service.PlatformOpenAI))
platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context())
require.True(t, ok)
require.Equal(t, service.PlatformOpenAI, platform)
}
func TestCompositeTargetPlatformAllowedRejectsWrongOrUnknownModel(t *testing.T) {
gin.SetMode(gin.TestMode)
for _, tc := range []struct {
name string
model string
}{
{name: "wrong provider", model: "claude-sonnet-4-5"},
{name: "unknown provider", model: "llama-4-maverick"},
} {
t.Run(tc.name, func(t *testing.T) {
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest("POST", "/v1/embeddings", nil)
apiKey := &service.APIKey{Group: &service.Group{Platform: service.PlatformComposite}}
require.False(t, compositeTargetPlatformAllowed(c, apiKey, tc.model, service.PlatformOpenAI))
})
}
}
+50 -1
View File
@@ -169,6 +169,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
body = parsedReq.Body.Bytes()
reqModel := parsedReq.Model
reqStream := parsedReq.Stream
ensureCompositeTargetPlatform(c, apiKey, reqModel)
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream))
// 解析渠道级模型映射
@@ -259,10 +260,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
zap.String("metadata_user_id_raw", parsedReq.MetadataUserID),
)
// 获取平台:优先使用强制平台(/antigravity 路由,中间件已设置 request.Context),否则使用分组平台
// 获取平台:优先使用强制平台(/antigravity 路由),其次使用 composite 解析出的目标平台,否则使用分组平台
platform := ""
if forcePlatform, ok := middleware2.GetForcePlatformFromContext(c); ok {
platform = forcePlatform
} else if resolvedPlatform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok {
platform = resolvedPlatform
} else if apiKey.Group != nil {
platform = apiKey.Group.Platform
}
@@ -1016,6 +1019,16 @@ func (h *GatewayHandler) Models(c *gin.Context) {
platform = forcedPlatform
}
if platform == service.PlatformComposite {
availableModels := h.compositeAvailableModels(c.Request.Context(), groupID)
if len(availableModels) > 0 {
writeModelsList(c, service.PlatformComposite, availableModels)
return
}
writeModelsList(c, service.PlatformComposite, defaultModelIDsForPlatform(service.PlatformComposite))
return
}
// Get available models from account configurations for the selected group platform.
availableModels := h.gatewayService.GetAvailableModels(c.Request.Context(), groupID, platform)
if apiKey != nil && apiKey.Group != nil && apiKey.Group.CustomModelsListEnabled() {
@@ -1057,6 +1070,28 @@ func (h *GatewayHandler) Models(c *gin.Context) {
})
}
func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID *int64) []string {
if h == nil || h.gatewayService == nil {
return nil
}
seen := make(map[string]struct{})
models := make([]string, 0)
for _, platform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformGrok} {
for _, model := range h.gatewayService.GetAvailableModels(ctx, groupID, platform) {
model = strings.TrimSpace(model)
if model == "" {
continue
}
if _, ok := seen[model]; ok {
continue
}
seen[model] = struct{}{}
models = append(models, model)
}
}
return models
}
func writeModelsList(c *gin.Context, platform string, modelIDs []string) {
if platform == service.PlatformGrok {
writeGrokModelsList(c, modelIDs)
@@ -1257,6 +1292,19 @@ func defaultModelIDsForPlatform(platform string) []string {
return mergeModelIDs(ids, nil)
case service.PlatformGrok:
return xai.DefaultModelIDs()
case service.PlatformComposite:
ids := make([]string, 0)
seen := make(map[string]struct{})
for _, concretePlatform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformGrok} {
for _, id := range defaultModelIDsForPlatform(concretePlatform) {
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
ids = append(ids, id)
}
}
return ids
default:
ids := make([]string, 0, len(claude.DefaultModels))
for _, model := range claude.DefaultModels {
@@ -1889,6 +1937,7 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) {
body = parsedReq.Body.Bytes()
// count_tokens 走 messages 严格校验时,复用已解析请求,避免二次反序列化。
SetClaudeCodeClientContext(c, body, parsedReq)
ensureCompositeTargetPlatform(c, apiKey, parsedReq.Model)
reqLog = reqLog.With(zap.String("model", parsedReq.Model), zap.Bool("stream", parsedReq.Stream))
// 在请求上下文中记录 thinking 状态,供 Antigravity 最终模型 key 推导/模型维度限流使用
c.Request = c.Request.WithContext(service.WithThinkingEnabled(c.Request.Context(), parsedReq.ThinkingEnabled, h.metadataBridgeEnabled()))
@@ -75,6 +75,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
return
}
reqModel := modelResult.String()
ensureCompositeTargetPlatform(c, apiKey, reqModel)
reqStream, ok := parseOpenAICompatibleStream(body)
if !ok {
h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage)
@@ -147,10 +148,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
APIKeyID: apiKey.ID,
}
sessionHash := h.gatewayService.GenerateSessionHash(parsedReq)
groupPlatform := ""
if apiKey.Group != nil {
groupPlatform = apiKey.Group.Platform
}
groupPlatform := effectiveAPIKeyPlatform(c, apiKey)
selectionSessionHash := sessionHash
if groupPlatform == service.PlatformGemini && selectionSessionHash != "" {
selectionSessionHash = "gemini:" + selectionSessionHash
@@ -75,6 +75,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
return
}
reqModel := modelResult.String()
ensureCompositeTargetPlatform(c, apiKey, reqModel)
reqStream, ok := parseOpenAICompatibleStream(body)
if !ok {
h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage)
@@ -85,7 +86,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
setOpsRequestContext(c, reqModel, reqStream)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
requestCtx := c.Request.Context()
if service.IsImageGenerationIntentForPlatform("/v1/responses", reqModel, body, openAICompatibleRequestPlatform(apiKey)) {
if service.IsImageGenerationIntentForPlatform("/v1/responses", reqModel, body, openAICompatibleRequestPlatform(c.Request.Context(), apiKey)) {
requestCtx = service.WithOpenAIImageGenerationIntent(requestCtx)
}
@@ -163,7 +164,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(requestCtx, apiKey.GroupID, sessionHash, reqModel, fs.FailedAccountIDs, "", int64(0))
if err != nil {
if len(fs.FailedAccountIDs) == 0 {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformAnthropic)
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, effectiveAPIKeyPlatform(c, apiKey))
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -41,7 +41,7 @@ func (h *GatewayHandler) GeminiV1BetaListModels(c *gin.Context) {
}
// 检查平台:优先使用强制平台(/antigravity 路由),否则要求 gemini 分组
forcePlatform, hasForcePlatform := middleware.GetForcePlatformFromContext(c)
if !hasForcePlatform && (apiKey.Group == nil || apiKey.Group.Platform != service.PlatformGemini) {
if !hasForcePlatform && effectiveAPIKeyPlatform(c, apiKey) != service.PlatformGemini {
googleError(c, http.StatusBadRequest, "API key group platform is not gemini")
return
}
@@ -88,7 +88,7 @@ func (h *GatewayHandler) GeminiV1BetaGetModel(c *gin.Context) {
}
// 检查平台:优先使用强制平台(/antigravity 路由),否则要求 gemini 分组
forcePlatform, hasForcePlatform := middleware.GetForcePlatformFromContext(c)
if !hasForcePlatform && (apiKey.Group == nil || apiKey.Group.Platform != service.PlatformGemini) {
if !hasForcePlatform && effectiveAPIKeyPlatform(c, apiKey) != service.PlatformGemini {
googleError(c, http.StatusBadRequest, "API key group platform is not gemini")
return
}
@@ -155,7 +155,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
// 检查平台:优先使用强制平台(/antigravity 路由,中间件已设置 request.Context),否则要求 gemini 分组
if !middleware.HasForcePlatform(c) {
if apiKey.Group == nil || apiKey.Group.Platform != service.PlatformGemini {
if effectiveAPIKeyPlatform(c, apiKey) != service.PlatformGemini {
googleError(c, http.StatusBadRequest, "API key group platform is not gemini")
return
}
+5 -1
View File
@@ -116,13 +116,17 @@ func classifyOpenAICompatibleNoAccountErrorFromGin(
routingModel string,
displayModel string,
) noAccountErrorClassification {
ctx := context.Background()
if c != nil && c.Request != nil {
ctx = c.Request.Context()
}
return classifyNoAccountErrorFromGin(
c,
diag,
apiKey,
routingModel,
displayModel,
openAICompatibleRequestPlatform(apiKey),
openAICompatibleRequestPlatform(ctx, apiKey),
)
}
@@ -75,6 +75,11 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
return
}
reqModel := modelResult.String()
ensureCompositeTargetPlatform(c, apiKey, reqModel)
if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups")
return
}
if apiKey.Group != nil && apiKey.Group.Platform == service.PlatformOpenAI {
if cappedBody, changed := service.ApplyOpenAIReasoningEffortPolicy(body, apiKey.Group.MaxReasoningEffort, apiKey.Group.ReasoningEffortMappings); changed {
body = cappedBody
@@ -111,7 +116,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
routingStart := time.Now()
@@ -71,6 +71,11 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
return
}
reqModel := modelResult.String()
ensureCompositeTargetPlatform(c, apiKey, reqModel)
if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups")
return
}
reqLog = reqLog.With(zap.String("model", reqModel))
setOpsRequestContext(c, reqModel, false)
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
@@ -116,6 +116,11 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
}
reqModel := parsedReq.Model
ensureCompositeTargetPlatform(c, apiKey, reqModel)
if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI) {
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups")
return
}
routingModel := service.NormalizeOpenAICompatRequestedModel(reqModel)
preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(apiKey, reqModel)
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", parsedReq.Stream))
@@ -155,11 +160,11 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
false,
false,
false,
openAICompatibleRequestPlatform(apiKey),
openAICompatibleRequestPlatform(c.Request.Context(), apiKey),
)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
if err != nil {
requestPlatform := openAICompatibleRequestPlatform(apiKey)
requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey)
reqLog.Warn("openai_count_tokens.account_select_failed", zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)))
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel)
if !cls.ModelNotFound {
@@ -117,7 +117,13 @@ func wrapUsageRecordTaskContext(parent context.Context, task service.UsageRecord
}
}
func openAICompatibleRequestPlatform(apiKey *service.APIKey) string {
func openAICompatibleRequestPlatform(ctx context.Context, apiKey *service.APIKey) string {
if platform, ok := service.ResolvedTargetPlatformFromContext(ctx); ok {
if platform == service.PlatformGrok {
return service.PlatformGrok
}
return service.PlatformOpenAI
}
if apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == service.PlatformGrok {
return service.PlatformGrok
}
@@ -253,6 +259,11 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
return
}
reqModel := modelResult.String()
ensureCompositeTargetPlatform(c, apiKey, reqModel)
if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI, service.PlatformGrok) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups")
return
}
if apiKey.Group != nil && apiKey.Group.Platform == service.PlatformOpenAI {
if cappedBody, changed := service.ApplyOpenAIReasoningEffortPolicy(body, apiKey.Group.MaxReasoningEffort, apiKey.Group.ReasoningEffortMappings); changed {
body = cappedBody
@@ -332,7 +343,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
// Get subscription info (may be nil)
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
routingStart := time.Now()
@@ -417,7 +428,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "compact_not_supported", "No available accounts support /responses/compact", streamStarted)
return
}
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, requestPlatform)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -432,7 +443,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
return
}
if selection == nil || selection.Account == nil {
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, requestPlatform)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
}
@@ -861,6 +872,11 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
return
}
reqModel := modelResult.String()
ensureCompositeTargetPlatform(c, apiKey, reqModel)
if !compositeTargetPlatformAllowed(c, apiKey, reqModel, service.PlatformOpenAI) {
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups")
return
}
routingModel := service.NormalizeOpenAICompatRequestedModel(reqModel)
preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(apiKey, reqModel)
reqStream := gjson.GetBytes(body, "stream").Bool()
@@ -885,7 +901,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
}
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
routingStart := time.Now()
@@ -1476,6 +1492,15 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "model is required in first response.create payload")
return
}
ensureCompositeTargetPlatform(c, apiKey, reqModel)
ctx = c.Request.Context()
if apiKey.Group != nil && apiKey.Group.Platform == service.PlatformComposite {
platform, ok := service.ResolvedTargetPlatformFromContext(ctx)
if !ok || (platform != service.PlatformOpenAI && platform != service.PlatformGrok) {
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "Responses WebSocket API only supports OpenAI-compatible models for composite groups")
return
}
}
previousResponseID := strings.TrimSpace(gjson.GetBytes(firstMessage, "previous_response_id").String())
previousResponseIDKind := service.ClassifyOpenAIPreviousResponseIDKind(previousResponseID)
if previousResponseID != "" && previousResponseIDKind == service.OpenAIPreviousResponseIDKindMessageID {
@@ -1567,7 +1592,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
}
subscription, _ := middleware2.GetSubscriptionFromContext(c)
requestPlatform := openAICompatibleRequestPlatform(apiKey)
requestPlatform := openAICompatibleRequestPlatform(ctx, apiKey)
requiredTransport := service.OpenAIUpstreamTransportResponsesWebsocketV2Ingress
if requestPlatform == service.PlatformGrok {
requiredTransport = service.OpenAIUpstreamTransportHTTPSSE
@@ -74,6 +74,11 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
return
}
requestModel := parsed.Model
ensureCompositeTargetPlatform(c, apiKey, requestModel)
if !compositeTargetPlatformAllowed(c, apiKey, requestModel, service.PlatformOpenAI) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by this OpenAI-compatible endpoint for composite groups")
return
}
reqLog = reqLog.With(
zap.String("model", requestModel),
+3
View File
@@ -8,6 +8,9 @@ const (
// ForcePlatform 强制平台(用于 /antigravity 路由),由 middleware.ForcePlatform 设置
ForcePlatform Key = "ctx_force_platform"
// ResolvedTargetPlatform 是 composite 分组按请求模型解析出的真实目标平台。
ResolvedTargetPlatform Key = "ctx_resolved_target_platform"
// RequestID 为服务端生成/透传的请求 ID。
RequestID Key = "ctx_request_id"
@@ -0,0 +1,46 @@
package routes
import (
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestCompositeTargetPlatformMiddlewareResolvesModelAndRestoresBody(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(gin.HandlerFunc(servermiddleware.APIKeyAuthMiddleware(func(c *gin.Context) {
groupID := int64(1)
c.Set(string(servermiddleware.ContextKeyAPIKey), &service.APIKey{
GroupID: &groupID,
Group: &service.Group{Platform: service.PlatformComposite},
})
c.Next()
})))
router.Use(compositeTargetPlatformMiddleware())
router.POST("/", func(c *gin.Context) {
platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context())
require.True(t, ok)
require.Equal(t, service.PlatformOpenAI, platform)
body, err := io.ReadAll(c.Request.Body)
require.NoError(t, err)
require.JSONEq(t, `{"model":"gpt-5"}`, string(body))
c.Status(http.StatusNoContent)
})
req := httptest.NewRequest(http.MethodPost, "/", strings.NewReader(`{"model":"gpt-5"}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNoContent, w.Code)
}
+92 -20
View File
@@ -1,14 +1,21 @@
package routes
import (
"bytes"
"errors"
"io"
"net/http"
"strconv"
"strings"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/handler"
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
)
// RegisterGatewayRoutes 注册 API 网关路由(Claude/OpenAI/Gemini 兼容)
@@ -27,6 +34,8 @@ func RegisterGatewayRoutes(
clientRequestID := middleware.ClientRequestID()
opsErrorLogger := handler.OpsErrorLoggerMiddleware(opsService)
endpointNorm := handler.InboundEndpointMiddleware()
compositeTarget := compositeTargetPlatformMiddleware()
compositeGeminiTarget := compositeImplicitTargetPlatformMiddleware(service.PlatformGemini)
// 未分组 Key 拦截中间件(按协议格式区分错误响应)
requireGroupAnthropic := middleware.RequireGroupAssignment(settingService, middleware.AnthropicErrorWriter)
@@ -60,6 +69,9 @@ func RegisterGatewayRoutes(
}
h.Gateway.Models(c)
}
isOpenAIOnlyEndpointGatewayPlatform := func(c *gin.Context) bool {
return getGroupPlatform(c) == service.PlatformOpenAI
}
imagesHandler := func(c *gin.Context) {
switch getGroupPlatform(c) {
case service.PlatformOpenAI:
@@ -139,6 +151,7 @@ func RegisterGatewayRoutes(
gateway.Use(endpointNorm)
gateway.Use(gin.HandlerFunc(apiKeyAuth))
gateway.GET("/sub2api/billing", h.Gateway.KeyBillingInfo)
gateway.Use(compositeTarget)
gateway.Use(requireGroupAnthropic)
{
// /v1/messages: auto-route based on group platform
@@ -185,7 +198,7 @@ func RegisterGatewayRoutes(
h.Gateway.ChatCompletions(c)
})
gateway.POST("/embeddings", textBodyLimit, func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
if !isOpenAIOnlyEndpointGatewayPlatform(c) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
@@ -226,6 +239,7 @@ func RegisterGatewayRoutes(
gemini.Use(opsErrorLogger)
gemini.Use(endpointNorm)
gemini.Use(middleware.APIKeyAuthWithSubscriptionGoogle(apiKeyService, subscriptionService, cfg))
gemini.Use(compositeGeminiTarget)
gemini.Use(requireGroupGoogle)
{
gemini.GET("/models", h.Gateway.GeminiV1BetaListModels)
@@ -242,16 +256,16 @@ func RegisterGatewayRoutes(
}
h.Gateway.Responses(c)
}
r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler)
r.POST("/alpha/search", textBodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.AlphaSearch)
r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, responsesHandler)
r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, responsesHandler)
r.POST("/alpha/search", textBodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, h.OpenAIGateway.AlphaSearch)
r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) {
h.OpenAIGateway.ResponsesWebSocket(c)
})
r.GET("/models", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, modelsHandler)
r.POST("/messages/count_tokens", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, countTokensHandler)
r.POST("/messages/count_tokens", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, countTokensHandler)
codexDirect := r.Group("/backend-api/codex")
codexDirect.Use(bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic)
codexDirect.Use(bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic)
{
codexDirect.POST("/responses", responsesHandler)
codexDirect.POST("/responses/*subpath", responsesHandler)
@@ -262,15 +276,15 @@ func RegisterGatewayRoutes(
codexDirect.GET("/models", h.OpenAIGateway.CodexModels)
}
// OpenAI Chat Completions API(不带v1前缀的别名)— auto-route based on group platform
r.POST("/chat/completions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
r.POST("/chat/completions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) {
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
h.OpenAIGateway.ChatCompletions(c)
return
}
h.Gateway.ChatCompletions(c)
})
r.POST("/embeddings", textBodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
r.POST("/embeddings", textBodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) {
if !isOpenAIOnlyEndpointGatewayPlatform(c) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
"error": gin.H{
@@ -282,16 +296,16 @@ func RegisterGatewayRoutes(
}
h.OpenAIGateway.Embeddings(c)
})
r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler)
r.POST("/images/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler)
r.POST("/images/generations/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.AsyncImage.Submit)
r.POST("/images/edits/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.AsyncImage.Submit)
r.GET("/images/tasks/:task_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.AsyncImage.Get)
r.POST("/videos/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoGenerationHandler)
r.POST("/videos/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoEditHandler)
r.POST("/videos/extensions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoExtensionHandler)
r.GET("/videos/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoStatusHandler)
r.GET("/videos/:request_id/content", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoContentHandler)
r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, imagesHandler)
r.POST("/images/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, imagesHandler)
r.POST("/images/generations/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, h.AsyncImage.Submit)
r.POST("/images/edits/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, h.AsyncImage.Submit)
r.GET("/images/tasks/:task_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, h.AsyncImage.Get)
r.POST("/videos/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoGenerationHandler)
r.POST("/videos/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoEditHandler)
r.POST("/videos/extensions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoExtensionHandler)
r.GET("/videos/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoStatusHandler)
r.GET("/videos/:request_id/content", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoContentHandler)
// Antigravity 模型列表
r.GET("/antigravity/models", gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.Gateway.AntigravityModels)
@@ -334,5 +348,63 @@ func getGroupPlatform(c *gin.Context) string {
if !ok || apiKey.Group == nil {
return ""
}
if apiKey.Group.Platform == service.PlatformComposite {
if platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context()); ok {
return platform
}
}
return apiKey.Group.Platform
}
func compositeTargetPlatformMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
apiKey, ok := middleware.GetAPIKeyFromContext(c)
if !ok || apiKey == nil || apiKey.Group == nil || apiKey.Group.Platform != service.PlatformComposite {
c.Next()
return
}
if c.Request == nil || c.Request.Method == http.MethodGet {
c.Next()
return
}
body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
if err != nil {
status := http.StatusBadRequest
message := "Failed to read request body"
var maxErr *http.MaxBytesError
if errors.As(err, &maxErr) {
status = http.StatusRequestEntityTooLarge
message = "Request body is too large"
}
c.JSON(status, gin.H{"error": gin.H{"type": "invalid_request_error", "message": message}})
c.Abort()
return
}
resetRequestBody(c, body)
model := strings.TrimSpace(gjson.GetBytes(body, "model").String())
if model != "" {
if platform, ok := service.DetectModelPlatform(model); ok {
c.Request = c.Request.WithContext(service.WithResolvedTargetPlatform(c.Request.Context(), platform))
}
}
c.Next()
}
}
func compositeImplicitTargetPlatformMiddleware(platform string) gin.HandlerFunc {
return func(c *gin.Context) {
apiKey, ok := middleware.GetAPIKeyFromContext(c)
if ok && apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == service.PlatformComposite {
c.Request = c.Request.WithContext(service.WithResolvedTargetPlatform(c.Request.Context(), platform))
}
c.Next()
}
}
func resetRequestBody(c *gin.Context, body []byte) {
c.Request.Body = io.NopCloser(bytes.NewReader(body))
c.Request.ContentLength = int64(len(body))
c.Request.Header.Set("Content-Length", strconv.Itoa(len(body)))
}
+22 -4
View File
@@ -115,6 +115,8 @@ func defaultModelsListCandidateIDs(platform string) []string {
return ids
case PlatformGrok:
return xai.DefaultModelIDs()
case PlatformComposite:
return nil
default:
ids := make([]string, 0, len(claude.DefaultModels))
for _, model := range claude.DefaultModels {
@@ -130,6 +132,22 @@ func defaultAllowImageGenerationForPlatform(platform string) bool {
return platform == PlatformGrok
}
func canCopyAccountsFromGroupPlatform(targetPlatform, sourcePlatform string) bool {
if targetPlatform == PlatformComposite {
return sourcePlatform == PlatformComposite || isConcreteRequestPlatform(sourcePlatform)
}
return sourcePlatform == targetPlatform
}
func groupSupportsOAuthOnlyFilter(platform string) bool {
return platform == PlatformOpenAI ||
platform == PlatformAntigravity ||
platform == PlatformAnthropic ||
platform == PlatformGemini ||
platform == PlatformGrok ||
platform == PlatformComposite
}
func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupInput) (*Group, error) {
if input.RateMultiplier <= 0 {
return nil, errors.New("rate_multiplier must be > 0")
@@ -255,7 +273,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn
if err != nil {
return nil, fmt.Errorf("source group %d not found: %w", srcGroupID, err)
}
if srcGroup.Platform != platform {
if !canCopyAccountsFromGroupPlatform(platform, srcGroup.Platform) {
return nil, fmt.Errorf("source group %d platform mismatch: expected %s, got %s", srcGroupID, platform, srcGroup.Platform)
}
}
@@ -321,7 +339,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn
}
// require_oauth_only: 过滤掉 apikey 类型账号
if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini || group.Platform == PlatformGrok) && len(accountIDsToCopy) > 0 {
if group.RequireOAuthOnly && groupSupportsOAuthOnlyFilter(group.Platform) && len(accountIDsToCopy) > 0 {
accounts, err := s.accountRepo.GetByIDs(ctx, accountIDsToCopy)
if err != nil {
return nil, fmt.Errorf("failed to fetch accounts for oauth filter: %w", err)
@@ -678,7 +696,7 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd
if err != nil {
return nil, fmt.Errorf("source group %d not found: %w", srcGroupID, err)
}
if srcGroup.Platform != group.Platform {
if !canCopyAccountsFromGroupPlatform(group.Platform, srcGroup.Platform) {
return nil, fmt.Errorf("source group %d platform mismatch: expected %s, got %s", srcGroupID, group.Platform, srcGroup.Platform)
}
}
@@ -695,7 +713,7 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd
}
// require_oauth_only: 过滤掉 apikey 类型账号
if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini || group.Platform == PlatformGrok) && len(accountIDsToCopy) > 0 {
if group.RequireOAuthOnly && groupSupportsOAuthOnlyFilter(group.Platform) && len(accountIDsToCopy) > 0 {
accounts, err := s.accountRepo.GetByIDs(ctx, accountIDsToCopy)
if err != nil {
return nil, fmt.Errorf("failed to fetch accounts for oauth filter: %w", err)
@@ -0,0 +1,114 @@
package service
import (
"context"
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
)
// WithResolvedTargetPlatform stores the concrete provider chosen for a request
// made through a composite group.
func WithResolvedTargetPlatform(ctx context.Context, platform string) context.Context {
platform = strings.TrimSpace(platform)
if ctx == nil || platform == "" {
return ctx
}
return context.WithValue(ctx, ctxkey.ResolvedTargetPlatform, platform)
}
// ResolvedTargetPlatformFromContext returns the concrete provider chosen for
// the current request, if one was resolved.
func ResolvedTargetPlatformFromContext(ctx context.Context) (string, bool) {
if ctx == nil {
return "", false
}
platform, ok := ctx.Value(ctxkey.ResolvedTargetPlatform).(string)
platform = strings.TrimSpace(platform)
if !ok || platform == "" {
return "", false
}
return platform, true
}
// DetectModelPlatform maps common public model IDs to the concrete provider
// platform used by sub2api. It intentionally returns false for ambiguous model
// names so composite groups fail closed instead of guessing.
func DetectModelPlatform(model string) (string, bool) {
normalized := strings.ToLower(strings.TrimSpace(model))
if normalized == "" {
return "", false
}
normalized = strings.TrimPrefix(normalized, "models/")
if slash := strings.IndexByte(normalized, '/'); slash > 0 {
provider := strings.TrimSpace(normalized[:slash])
rest := strings.TrimSpace(normalized[slash+1:])
switch provider {
case "anthropic", "claude":
return PlatformAnthropic, true
case "openai", "chatgpt":
return PlatformOpenAI, true
case "google", "google-ai-studio", "gemini":
return PlatformGemini, true
case "xai", "x-ai", "grok":
return PlatformGrok, true
}
if rest != "" {
normalized = strings.TrimPrefix(rest, "models/")
}
}
switch {
case strings.HasPrefix(normalized, "anthropic.claude-"),
strings.HasPrefix(normalized, "claude-"):
return PlatformAnthropic, true
case strings.HasPrefix(normalized, "gpt-"),
strings.HasPrefix(normalized, "chatgpt-"),
strings.HasPrefix(normalized, "codex-"),
strings.HasPrefix(normalized, "text-embedding-"),
strings.HasPrefix(normalized, "text-moderation-"),
strings.HasPrefix(normalized, "omni-moderation-"),
strings.HasPrefix(normalized, "dall-e-"),
strings.HasPrefix(normalized, "gpt-image-"),
strings.HasPrefix(normalized, "tts-"),
strings.HasPrefix(normalized, "whisper-"),
hasOpenAISeriesPrefix(normalized):
return PlatformOpenAI, true
case strings.HasPrefix(normalized, "gemini-"),
strings.HasPrefix(normalized, "learnlm-"):
return PlatformGemini, true
case normalized == "grok" || strings.HasPrefix(normalized, "grok-"):
return PlatformGrok, true
default:
return "", false
}
}
func hasOpenAISeriesPrefix(model string) bool {
for _, prefix := range []string{"o1", "o3", "o4", "o5"} {
if model == prefix || strings.HasPrefix(model, prefix+"-") {
return true
}
}
return false
}
func resolveCompositeTargetPlatform(ctx context.Context, group *Group, requestedModel string) (string, bool) {
if platform, ok := ResolvedTargetPlatformFromContext(ctx); ok {
return platform, true
}
if group == nil || group.Platform != PlatformComposite {
return "", false
}
return DetectModelPlatform(requestedModel)
}
func isConcreteRequestPlatform(platform string) bool {
switch platform {
case PlatformAnthropic, PlatformOpenAI, PlatformGemini, PlatformAntigravity, PlatformGrok:
return true
default:
return false
}
}
@@ -0,0 +1,58 @@
package service
import (
"context"
"testing"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
"github.com/stretchr/testify/require"
)
func TestDetectModelPlatform(t *testing.T) {
tests := []struct {
name string
model string
platform string
ok bool
}{
{name: "claude", model: "claude-sonnet-4-5", platform: PlatformAnthropic, ok: true},
{name: "anthropic prefix", model: "anthropic/claude-opus-4-5", platform: PlatformAnthropic, ok: true},
{name: "gpt", model: "gpt-5.1", platform: PlatformOpenAI, ok: true},
{name: "o series", model: "o3-mini", platform: PlatformOpenAI, ok: true},
{name: "embedding", model: "text-embedding-3-large", platform: PlatformOpenAI, ok: true},
{name: "gemini", model: "gemini-3-pro", platform: PlatformGemini, ok: true},
{name: "gemini models prefix", model: "models/gemini-2.5-flash", platform: PlatformGemini, ok: true},
{name: "learnlm", model: "learnlm-2.0-flash-experimental", platform: PlatformGemini, ok: true},
{name: "grok", model: "grok-4", platform: PlatformGrok, ok: true},
{name: "xai prefix", model: "xai/grok-4", platform: PlatformGrok, ok: true},
{name: "unknown", model: "llama-4-maverick", ok: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
platform, ok := DetectModelPlatform(tt.model)
require.Equal(t, tt.ok, ok)
require.Equal(t, tt.platform, platform)
})
}
}
func TestQuotaPlatformCompositeUsesResolvedOrForceOnly(t *testing.T) {
apiKey := &APIKey{Group: &Group{Platform: PlatformComposite}}
require.Equal(t, "", QuotaPlatform(context.Background(), apiKey))
require.Equal(t, PlatformGemini, QuotaPlatform(WithResolvedTargetPlatform(context.Background(), PlatformGemini), apiKey))
require.Equal(t, PlatformAntigravity, QuotaPlatform(context.WithValue(context.Background(), ctxkey.ForcePlatform, PlatformAntigravity), apiKey))
ctx := WithResolvedTargetPlatform(context.Background(), PlatformAnthropic)
ctx = context.WithValue(ctx, ctxkey.ForcePlatform, PlatformAntigravity)
require.Equal(t, PlatformAntigravity, QuotaPlatform(ctx, apiKey))
}
func TestSchedulerPlatformsForCompositeGroup(t *testing.T) {
require.ElementsMatch(t,
[]string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok},
schedulerPlatformsForGroup(PlatformComposite),
)
require.Equal(t, []string{PlatformAnthropic}, schedulerPlatformsForGroup(PlatformAnthropic))
}
@@ -43,6 +43,7 @@ const (
PlatformGemini = domain.PlatformGemini
PlatformAntigravity = domain.PlatformAntigravity
PlatformGrok = domain.PlatformGrok
PlatformComposite = domain.PlatformComposite
)
// AllowedQuotaPlatforms 是允许设置 user × platform quota 的平台列表(单一权威来源)。
+33 -6
View File
@@ -45,6 +45,14 @@ func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context
groupID = resolvedGroupID
ctx = s.withGroupContext(ctx, group)
platform = group.Platform
if group != nil && group.Platform == PlatformComposite {
targetPlatform, ok := resolveCompositeTargetPlatform(ctx, group, requestedModel)
if !ok {
return nil, fmt.Errorf("%w supporting model: %s (composite target platform unknown)", ErrNoAvailableAccounts, requestedModel)
}
platform = targetPlatform
ctx = WithResolvedTargetPlatform(ctx, targetPlatform)
}
} else {
// 无分组时只使用原生 anthropic 平台
platform = PlatformAnthropic
@@ -194,7 +202,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
}
}
platform, hasForcePlatform, err := s.resolvePlatform(ctx, groupID, group)
platform, hasForcePlatform, err := s.resolvePlatform(ctx, groupID, group, requestedModel)
if err != nil {
return nil, err
}
@@ -226,9 +234,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
return excluded
}
// 获取模型路由配置(仅 anthropic 平台)
// 获取模型路由配置(anthropic 目标平台;composite 分组按目标平台判断)
var routingAccountIDs []int64
if group != nil && requestedModel != "" && group.Platform == PlatformAnthropic {
if group != nil && requestedModel != "" && platform == PlatformAnthropic &&
(group.Platform == PlatformAnthropic || group.Platform == PlatformComposite) {
routingAccountIDs = group.GetRoutingAccountIDs(requestedModel)
if s.debugModelRoutingEnabled() {
logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] context group routing: group_id=%d model=%s enabled=%v rules=%d matched_ids=%v session=%s sticky_account=%d",
@@ -822,8 +831,9 @@ func (s *GatewayService) routingAccountIDsForRequest(ctx context.Context, groupI
}
return nil
}
// Preserve existing behavior: model routing only applies to anthropic groups.
if group.Platform != PlatformAnthropic {
// Model routing applies only to requests resolved to Anthropic. Composite
// groups may still use those rules once their model resolved to Anthropic.
if group.Platform != PlatformAnthropic && group.Platform != PlatformComposite {
if s.debugModelRoutingEnabled() {
logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] skip: non-anthropic group platform: group_id=%d group_platform=%s model=%s", group.ID, group.Platform, requestedModel)
}
@@ -888,12 +898,22 @@ func (s *GatewayService) checkClaudeCodeRestriction(ctx context.Context, groupID
return group, resolvedID, nil
}
func (s *GatewayService) resolvePlatform(ctx context.Context, groupID *int64, group *Group) (string, bool, error) {
func (s *GatewayService) resolvePlatform(ctx context.Context, groupID *int64, group *Group, requestedModel string) (string, bool, error) {
forcePlatform, hasForcePlatform := ctx.Value(ctxkey.ForcePlatform).(string)
if hasForcePlatform && forcePlatform != "" {
return forcePlatform, true, nil
}
if platform, ok := ResolvedTargetPlatformFromContext(ctx); ok {
return platform, false, nil
}
if group != nil {
if group.Platform == PlatformComposite {
targetPlatform, ok := resolveCompositeTargetPlatform(ctx, group, requestedModel)
if !ok {
return "", false, fmt.Errorf("%w supporting model: %s (composite target platform unknown)", ErrNoAvailableAccounts, requestedModel)
}
return targetPlatform, false, nil
}
return group.Platform, false, nil
}
if groupID != nil {
@@ -901,6 +921,13 @@ func (s *GatewayService) resolvePlatform(ctx context.Context, groupID *int64, gr
if err != nil {
return "", false, err
}
if group.Platform == PlatformComposite {
targetPlatform, ok := resolveCompositeTargetPlatform(ctx, group, requestedModel)
if !ok {
return "", false, fmt.Errorf("%w supporting model: %s (composite target platform unknown)", ErrNoAvailableAccounts, requestedModel)
}
return targetPlatform, false, nil
}
return group.Platform, false, nil
}
return PlatformAnthropic, false, nil
@@ -100,10 +100,19 @@ func PlatformFromAPIKey(apiKey *APIKey) string {
// 后扣运行在 worker 池的 background ctx 上没有 ForcePlatform,因此后扣平台由 handler
// 预先算定、经 RecordUsageInput.QuotaPlatform 传入,不要在后扣链路用 worker ctx 调用本函数。
func QuotaPlatform(ctx context.Context, apiKey *APIKey) string {
if fp, ok := ctx.Value(ctxkey.ForcePlatform).(string); ok && fp != "" {
return fp
if ctx != nil {
if fp, ok := ctx.Value(ctxkey.ForcePlatform).(string); ok && fp != "" {
return fp
}
}
return PlatformFromAPIKey(apiKey)
if platform, ok := ResolvedTargetPlatformFromContext(ctx); ok {
return platform
}
platform := PlatformFromAPIKey(apiKey)
if platform == PlatformComposite {
return ""
}
return platform
}
func (p *postUsageBillingParams) shouldDeductAPIKeyQuota() bool {
@@ -733,6 +742,9 @@ func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsage
quotaPlatform := input.QuotaPlatform
if quotaPlatform == "" {
quotaPlatform = PlatformFromAPIKey(apiKey)
if quotaPlatform == PlatformComposite && account != nil {
quotaPlatform = account.Platform
}
}
requestID := usageLog.RequestID
_, billingErr := applyUsageBilling(ctx, requestID, usageLog, &postUsageBillingParams{
@@ -94,11 +94,11 @@ const filteredGroups = computed(() => {
// antigravity 账户启用混合调度后,可选择 anthropic/gemini 分组
if (props.platform === 'antigravity' && props.mixedScheduling) {
result = result.filter(
(g) => g.platform === 'antigravity' || g.platform === 'anthropic' || g.platform === 'gemini'
(g) => g.platform === 'antigravity' || g.platform === 'anthropic' || g.platform === 'gemini' || g.platform === 'composite'
)
} else {
// 默认:只能选择同 platform 的分组
result = result.filter((g) => g.platform === props.platform)
// 默认:只能选择同 platform 的分组;composite 分组可接收任意具体平台账号
result = result.filter((g) => g.platform === props.platform || g.platform === 'composite')
}
}
if (isSearchable.value && searchText.value) {
@@ -25,6 +25,13 @@
d="M9.27 15.29l7.978-5.897c.391-.29.95-.177 1.137.272.98 2.369.542 5.215-1.41 7.169-1.951 1.954-4.667 2.382-7.149 1.406l-2.711 1.257c3.889 2.661 8.611 2.003 11.562-.953 2.341-2.344 3.066-5.539 2.388-8.42l.006.007c-.983-4.232.242-5.924 2.75-9.383.06-.082.12-.164.179-.248l-3.301 3.305v-.01L9.267 15.292M7.623 16.723c-2.792-2.67-2.31-6.801.071-9.184 1.761-1.763 4.647-2.483 7.166-1.425l2.705-1.25a7.808 7.808 0 00-1.829-1A8.975 8.975 0 005.984 5.83c-2.533 2.536-3.33 6.436-1.962 9.764 1.022 2.487-.653 4.246-2.34 6.022-.599.63-1.199 1.259-1.682 1.925l7.62-6.815"
/>
</svg>
<!-- Composite group icon -->
<svg v-else-if="platform === 'composite'" :class="sizeClass" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2">
<circle cx="6" cy="12" r="3" />
<circle cx="18" cy="6" r="3" />
<circle cx="18" cy="18" r="3" />
<path stroke-linecap="round" stroke-linejoin="round" d="M8.7 10.7 15.3 7.3M8.7 13.3l6.6 3.4" />
</svg>
<!-- Fallback: generic platform icon -->
<svg v-else :class="sizeClass" fill="currentColor" viewBox="0 0 24 24">
<path
@@ -944,6 +944,7 @@ export default {
gemini: 'Gemini',
antigravity: 'Antigravity',
grok: 'Grok',
composite: 'Composite',
},
deleteConfirm:
"Are you sure you want to delete '{name}'? All associated API keys will no longer belong to any group.",
@@ -877,6 +877,7 @@ export default {
gemini: 'Gemini',
antigravity: 'Antigravity',
grok: 'Grok',
composite: 'Composite',
},
saving: '保存中...',
noGroups: '暂无分组',
+1 -1
View File
@@ -492,7 +492,7 @@ export interface PaginationConfig {
// ==================== API Key & Group Types ====================
export type GroupPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' | 'grok'
export type GroupPlatform = 'anthropic' | 'openai' | 'gemini' | 'antigravity' | 'grok' | 'composite'
export type SubscriptionType = 'standard' | 'subscription'
+14 -2
View File
@@ -5,7 +5,7 @@
* instead of defining their own color mappings.
*/
export type Platform = 'anthropic' | 'openai' | 'antigravity' | 'gemini' | 'grok'
export type Platform = 'anthropic' | 'openai' | 'antigravity' | 'gemini' | 'grok' | 'composite'
// ── Badge (bg + text + border, for inline badges with border) ───────
const BADGE: Record<Platform, string> = {
@@ -14,6 +14,7 @@ const BADGE: Record<Platform, string> = {
antigravity: 'bg-purple-500/10 text-purple-600 border-purple-500/30 dark:text-purple-400',
gemini: 'bg-blue-500/10 text-blue-600 border-blue-500/30 dark:text-blue-400',
grok: 'bg-zinc-800/10 text-zinc-800 border-zinc-800/30 dark:bg-zinc-500/10 dark:text-zinc-200 dark:border-zinc-500/30',
composite: 'bg-cyan-500/10 text-cyan-700 border-cyan-500/30 dark:text-cyan-300',
}
const BADGE_DEFAULT = 'bg-slate-500/10 text-slate-600 border-slate-500/30 dark:text-slate-400'
@@ -24,6 +25,7 @@ const BADGE_LIGHT: Record<Platform, string> = {
antigravity: 'bg-purple-500/10 text-purple-600 dark:bg-purple-500/10 dark:text-purple-300',
gemini: 'bg-blue-500/10 text-blue-600 dark:bg-blue-500/10 dark:text-blue-300',
grok: 'bg-zinc-800/10 text-zinc-800 dark:bg-zinc-500/10 dark:text-zinc-200',
composite: 'bg-cyan-500/10 text-cyan-700 dark:bg-cyan-500/10 dark:text-cyan-300',
}
// ── Border ──────────────────────────────────────────────────────────
@@ -33,6 +35,7 @@ const BORDER: Record<Platform, string> = {
antigravity: 'border-purple-500/20 dark:border-purple-500/20',
gemini: 'border-blue-500/20 dark:border-blue-500/20',
grok: 'border-zinc-800/20 dark:border-zinc-500/20',
composite: 'border-cyan-500/20 dark:border-cyan-500/20',
}
const BORDER_DEFAULT = 'border-gray-200 dark:border-dark-700'
@@ -43,6 +46,7 @@ const ACCENT_BAR: Record<Platform, string> = {
antigravity: 'bg-gradient-to-r from-purple-400 to-purple-500',
gemini: 'bg-gradient-to-r from-blue-400 to-blue-500',
grok: 'bg-gradient-to-r from-zinc-700 to-zinc-900',
composite: 'bg-gradient-to-r from-slate-500 to-cyan-500',
}
const ACCENT_BAR_DEFAULT = 'bg-gradient-to-r from-primary-400 to-primary-500'
@@ -53,6 +57,7 @@ const TEXT: Record<Platform, string> = {
antigravity: 'text-purple-600 dark:text-purple-400',
gemini: 'text-blue-600 dark:text-blue-400',
grok: 'text-zinc-800 dark:text-zinc-200',
composite: 'text-cyan-700 dark:text-cyan-300',
}
const TEXT_DEFAULT = 'text-primary-600 dark:text-primary-400'
@@ -63,6 +68,7 @@ const ICON: Record<Platform, string> = {
antigravity: 'text-purple-500 dark:text-purple-400',
gemini: 'text-blue-500 dark:text-blue-400',
grok: 'text-zinc-800 dark:text-zinc-200',
composite: 'text-cyan-600 dark:text-cyan-300',
}
const ICON_DEFAULT = 'text-primary-500 dark:text-primary-400'
@@ -73,6 +79,7 @@ const BUTTON: Record<Platform, string> = {
antigravity: 'bg-purple-500 text-white hover:bg-purple-600 active:bg-purple-700 dark:bg-purple-500/80 dark:hover:bg-purple-500',
gemini: 'bg-blue-500 text-white hover:bg-blue-600 active:bg-blue-700 dark:bg-blue-500/80 dark:hover:bg-blue-500',
grok: 'bg-zinc-800 text-white hover:bg-zinc-900 active:bg-black dark:bg-zinc-700 dark:hover:bg-zinc-600',
composite: 'bg-cyan-700 text-white hover:bg-cyan-800 active:bg-cyan-900 dark:bg-cyan-600 dark:hover:bg-cyan-500',
}
const BUTTON_DEFAULT = 'bg-primary-500 text-white hover:bg-primary-600 dark:bg-primary-600 dark:hover:bg-primary-500'
@@ -83,6 +90,7 @@ const DISCOUNT: Record<Platform, string> = {
antigravity: 'bg-purple-100 text-purple-700 dark:bg-purple-900/40 dark:text-purple-300',
gemini: 'bg-blue-100 text-blue-700 dark:bg-blue-900/40 dark:text-blue-300',
grok: 'bg-zinc-100 text-zinc-800 dark:bg-zinc-800 dark:text-zinc-200',
composite: 'bg-cyan-100 text-cyan-800 dark:bg-cyan-900/40 dark:text-cyan-300',
}
const DISCOUNT_DEFAULT = 'bg-red-100 text-red-700 dark:bg-red-900/40 dark:text-red-300'
@@ -93,6 +101,7 @@ const GRADIENT: Record<Platform, string> = {
antigravity: 'from-purple-500 to-purple-600',
gemini: 'from-blue-500 to-blue-600',
grok: 'from-zinc-700 to-zinc-900',
composite: 'from-slate-600 to-cyan-600',
}
const GRADIENT_DEFAULT = 'from-primary-500 to-primary-600'
@@ -103,6 +112,7 @@ const GRADIENT_TEXT: Record<Platform, string> = {
antigravity: 'text-purple-100',
gemini: 'text-blue-100',
grok: 'text-zinc-100',
composite: 'text-cyan-100',
}
const GRADIENT_TEXT_DEFAULT = 'text-primary-100'
@@ -112,13 +122,14 @@ const GRADIENT_SUBTEXT: Record<Platform, string> = {
antigravity: 'text-purple-200',
gemini: 'text-blue-200',
grok: 'text-zinc-300',
composite: 'text-cyan-200',
}
const GRADIENT_SUBTEXT_DEFAULT = 'text-primary-200'
// ── Public API ──────────────────────────────────────────────────────
function isPlatform(p: string): p is Platform {
return p === 'anthropic' || p === 'openai' || p === 'antigravity' || p === 'gemini' || p === 'grok'
return p === 'anthropic' || p === 'openai' || p === 'antigravity' || p === 'gemini' || p === 'grok' || p === 'composite'
}
export function platformBadgeClass(p: string): string {
@@ -172,6 +183,7 @@ export function platformLabel(p: string): string {
case 'antigravity': return 'Antigravity'
case 'gemini': return 'Gemini'
case 'grok': return 'Grok'
case 'composite': return 'Composite'
default: return p || 'API'
}
}
+19 -6
View File
@@ -3841,6 +3841,7 @@ const platformOptions = computed(() => [
{ value: "gemini", label: "Gemini" },
{ value: "antigravity", label: "Antigravity" },
{ value: "grok", label: "Grok" },
{ value: "composite", label: "Composite" },
]);
const platformFilterOptions = computed(() => [
@@ -3850,6 +3851,7 @@ const platformFilterOptions = computed(() => [
{ value: "gemini", label: "Gemini" },
{ value: "antigravity", label: "Antigravity" },
{ value: "grok", label: "Grok" },
{ value: "composite", label: "Composite" },
]);
const editStatusOptions = computed(() => [
@@ -3936,29 +3938,40 @@ const invalidRequestFallbackOptionsForEdit = computed(() => {
return options;
});
// 复制账号的源分组选项(创建时)- 仅包含相同平台且有账号的分组
const canCopyAccountsFromGroup = (targetPlatform: GroupPlatform, sourcePlatform: GroupPlatform) =>
targetPlatform === "composite" || sourcePlatform === targetPlatform;
const copyAccountsGroupLabel = (g: AdminGroup) => {
const count = g.account_count || 0;
const platform = t("admin.groups.platforms." + g.platform);
return `${g.name} - ${platform} (${t("admin.groups.accountsCount", { count })})`;
};
// 复制账号的源分组选项(创建时)- 相同平台;composite 分组可汇总各平台账号
const copyAccountsGroupOptions = computed(() => {
const eligibleGroups = groups.value.filter(
(g) => g.platform === createForm.platform && (g.account_count || 0) > 0,
(g) =>
canCopyAccountsFromGroup(createForm.platform, g.platform) &&
(g.account_count || 0) > 0,
);
return eligibleGroups.map((g) => ({
value: g.id,
label: `${g.name} (${t("admin.groups.accountsCount", { count: g.account_count || 0 })})`,
label: copyAccountsGroupLabel(g),
}));
});
// 复制账号的源分组选项(编辑时)- 仅包含相同平台且有账号的分组,排除自身
// 复制账号的源分组选项(编辑时)- 相同平台;composite 分组可汇总各平台账号,排除自身
const copyAccountsGroupOptionsForEdit = computed(() => {
const currentId = editingGroup.value?.id;
const eligibleGroups = groups.value.filter(
(g) =>
g.platform === editForm.platform &&
canCopyAccountsFromGroup(editForm.platform, g.platform) &&
(g.account_count || 0) > 0 &&
g.id !== currentId,
);
return eligibleGroups.map((g) => ({
value: g.id,
label: `${g.name} (${t("admin.groups.accountsCount", { count: g.account_count || 0 })})`,
label: copyAccountsGroupLabel(g),
}));
});