mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 17:08:33 +08:00
Add composite group routing
This commit is contained in:
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)))
|
||||
}
|
||||
|
||||
@@ -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 的平台列表(单一权威来源)。
|
||||
|
||||
@@ -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: '暂无分组',
|
||||
|
||||
@@ -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'
|
||||
|
||||
|
||||
@@ -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'
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
}));
|
||||
});
|
||||
|
||||
|
||||
Reference in New Issue
Block a user