mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
fix: harden grok media routing
This commit is contained in:
@@ -14,7 +14,6 @@ import (
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/tidwall/gjson"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
@@ -84,7 +83,8 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
}
|
||||
|
||||
contentType := c.GetHeader("Content-Type")
|
||||
requestModel := service.ExtractGrokMediaModel(contentType, body)
|
||||
requestInfo := service.ParseGrokMediaRequest(contentType, body)
|
||||
requestModel := requestInfo.Model
|
||||
if endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) == "" {
|
||||
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
|
||||
return
|
||||
@@ -103,7 +103,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
|
||||
return
|
||||
}
|
||||
if moderationBody := grokMediaModerationBody(body); len(moderationBody) > 0 {
|
||||
if moderationBody := requestInfo.ModerationBody(); len(moderationBody) > 0 {
|
||||
decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, moderationBody)
|
||||
if decision != nil && decision.Blocked {
|
||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
@@ -149,6 +149,9 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
sessionSeed = []byte(requestID)
|
||||
}
|
||||
sessionHash := h.gatewayService.GenerateExplicitSessionHash(c, sessionSeed)
|
||||
if endpoint == service.GrokMediaEndpointVideoStatus {
|
||||
sessionHash = service.GrokMediaVideoRequestSessionHash(requestID)
|
||||
}
|
||||
requestCtx := c.Request.Context()
|
||||
failedAccountIDs := make(map[int64]struct{})
|
||||
sameAccountRetryCount := make(map[int64]int)
|
||||
@@ -294,6 +297,15 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
}
|
||||
|
||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil)
|
||||
if endpoint == service.GrokMediaEndpointVideosGenerations && strings.TrimSpace(result.ResponseID) != "" {
|
||||
if err := h.gatewayService.BindGrokMediaVideoRequestAccount(requestCtx, apiKey.GroupID, result.ResponseID, account.ID); err != nil {
|
||||
reqLog.Warn("grok_media.bind_video_request_account_failed",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.String("request_id", result.ResponseID),
|
||||
zap.Error(err),
|
||||
)
|
||||
}
|
||||
}
|
||||
if shouldRecordGrokMediaUsage(endpoint, requestModel) {
|
||||
recordGrokMediaUsage(c, h, reqLog, apiKey, subject, subscription, account, result, requestModel, body, requestID)
|
||||
}
|
||||
@@ -305,13 +317,6 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
}
|
||||
}
|
||||
|
||||
func grokMediaModerationBody(body []byte) []byte {
|
||||
if gjson.ValidBytes(body) {
|
||||
return body
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func shouldRecordGrokMediaUsage(endpoint service.GrokMediaEndpoint, requestModel string) bool {
|
||||
return endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) != ""
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user