Merge upstream/main into fix/api-double-billing

This commit is contained in:
benjamin
2026-07-13 11:10:33 +08:00
120 changed files with 6391 additions and 402 deletions
@@ -98,7 +98,7 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) {
require.Contains(t, rec.Body.String(), `"source":"active_probe"`)
require.Contains(t, rec.Body.String(), `"headers_observed":true`)
require.NotContains(t, rec.Body.String(), "access-token")
require.Equal(t, xai.DefaultBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Contains(t, string(upstream.lastBody), `"store":false`)
require.NotNil(t, repo.updates[42])
@@ -110,6 +110,7 @@ type CreateGroupRequest struct {
VideoPrice480P *float64 `json:"video_price_480p"`
VideoPrice720P *float64 `json:"video_price_720p"`
VideoPrice1080P *float64 `json:"video_price_1080p"`
WebSearchPricePerCall *float64 `json:"web_search_price_per_call"`
ClaudeCodeOnly bool `json:"claude_code_only"`
FallbackGroupID *int64 `json:"fallback_group_id"`
FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"`
@@ -163,6 +164,7 @@ type UpdateGroupRequest struct {
VideoPrice480P *float64 `json:"video_price_480p"`
VideoPrice720P *float64 `json:"video_price_720p"`
VideoPrice1080P *float64 `json:"video_price_1080p"`
WebSearchPricePerCall *float64 `json:"web_search_price_per_call"`
ClaudeCodeOnly *bool `json:"claude_code_only"`
FallbackGroupID *int64 `json:"fallback_group_id"`
FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"`
@@ -334,6 +336,7 @@ func (h *GroupHandler) Create(c *gin.Context) {
VideoPrice480P: req.VideoPrice480P,
VideoPrice720P: req.VideoPrice720P,
VideoPrice1080P: req.VideoPrice1080P,
WebSearchPricePerCall: req.WebSearchPricePerCall,
ClaudeCodeOnly: req.ClaudeCodeOnly,
FallbackGroupID: req.FallbackGroupID,
FallbackGroupIDOnInvalidRequest: req.FallbackGroupIDOnInvalidRequest,
@@ -402,6 +405,7 @@ func (h *GroupHandler) Update(c *gin.Context) {
VideoPrice480P: req.VideoPrice480P,
VideoPrice720P: req.VideoPrice720P,
VideoPrice1080P: req.VideoPrice1080P,
WebSearchPricePerCall: req.WebSearchPricePerCall,
ClaudeCodeOnly: req.ClaudeCodeOnly,
FallbackGroupID: req.FallbackGroupID,
FallbackGroupIDOnInvalidRequest: req.FallbackGroupIDOnInvalidRequest,
+1
View File
@@ -199,6 +199,7 @@ func groupFromServiceBase(g *service.Group) Group {
VideoPrice480P: g.VideoPrice480P,
VideoPrice720P: g.VideoPrice720P,
VideoPrice1080P: g.VideoPrice1080P,
WebSearchPricePerCall: g.WebSearchPricePerCall,
ClaudeCodeOnly: g.ClaudeCodeOnly,
FallbackGroupID: g.FallbackGroupID,
FallbackGroupIDOnInvalidRequest: g.FallbackGroupIDOnInvalidRequest,
+2
View File
@@ -120,6 +120,8 @@ type Group struct {
VideoPrice480P *float64 `json:"video_price_480p"`
VideoPrice720P *float64 `json:"video_price_720p"`
VideoPrice1080P *float64 `json:"video_price_1080p"`
// Codex alpha/search 网页搜索单次价格(USD/次);null 表示使用默认价 0.01
WebSearchPricePerCall *float64 `json:"web_search_price_per_call"`
// Claude Code 客户端限制
ClaudeCodeOnly bool `json:"claude_code_only"`
+9 -3
View File
@@ -18,6 +18,7 @@ const (
EndpointMessages = "/v1/messages"
EndpointChatCompletions = "/v1/chat/completions"
EndpointEmbeddings = "/v1/embeddings"
EndpointAlphaSearch = "/v1/alpha/search"
EndpointResponses = "/v1/responses"
EndpointResponsesCompact = "/v1/responses/compact"
EndpointImagesGenerations = "/v1/images/generations"
@@ -75,6 +76,8 @@ func NormalizeInboundEndpoint(path string) string {
switch {
case strings.Contains(path, EndpointEmbeddings):
return EndpointEmbeddings
case strings.Contains(path, EndpointAlphaSearch) || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/alpha/search") || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/backend-api/codex/alpha/search"):
return EndpointAlphaSearch
case strings.Contains(path, EndpointChatCompletions):
return EndpointChatCompletions
case strings.Contains(path, EndpointMessages):
@@ -155,8 +158,11 @@ func isBareOrSubpathOf(path, root string) bool {
// account platform and the normalized inbound endpoint.
//
// Platform-specific rules:
// - OpenAI always forwards to /v1/responses (with optional subpath
// such as /v1/responses/compact preserved from the raw URL).
// - OpenAI and Grok text compatibility routes forward to /v1/responses
// (with optional subpath such as /v1/responses/compact preserved from
// the raw URL); native endpoints such as embeddings and alpha search
// retain their paths. Grok raw Chat requests override this through the
// forwarding result consumed by resolveOpenAIUpstreamEndpoint.
// - Anthropic → /v1/messages
// - Gemini → /v1beta/models
// - Antigravity → /v1/messages (Claude) or gemini (Gemini)
@@ -167,7 +173,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
switch platform {
case service.PlatformOpenAI, service.PlatformGrok:
if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos {
if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos {
return inbound
}
// OpenAI forwards everything to the Responses API.
+59
View File
@@ -25,6 +25,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
{"/v1/messages", EndpointMessages},
{"/v1/chat/completions", EndpointChatCompletions},
{"/v1/embeddings", EndpointEmbeddings},
{"/v1/alpha/search", EndpointAlphaSearch},
{"/v1/responses", EndpointResponses},
{"/v1/responses/compact", EndpointResponsesCompact},
{"/v1/responses/compact/detail", EndpointResponsesCompact},
@@ -50,11 +51,13 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
{"/responses", EndpointResponses},
{"/responses/compact", EndpointResponsesCompact},
{"/responses/compact/detail", EndpointResponsesCompact},
{"/alpha/search", EndpointAlphaSearch},
// Bare Codex direct alias route — root vs. compact.
{"/backend-api/codex/responses", EndpointResponses},
{"/backend-api/codex/responses/compact", EndpointResponsesCompact},
{"/backend-api/codex/responses/compact/detail", EndpointResponsesCompact},
{"/backend-api/codex/alpha/search", EndpointAlphaSearch},
// Must NOT generalize to arbitrary paths merely ending in
// "/responses" (or "/responses/compact") that are unrelated to
@@ -119,8 +122,11 @@ func TestDeriveUpstreamEndpoint(t *testing.T) {
{"openai from messages", EndpointMessages, "/v1/messages", service.PlatformOpenAI, EndpointResponses},
{"openai from completions", EndpointChatCompletions, "/v1/chat/completions", service.PlatformOpenAI, EndpointResponses},
{"openai embeddings", EndpointEmbeddings, "/v1/embeddings", service.PlatformOpenAI, EndpointEmbeddings},
{"openai alpha search", EndpointAlphaSearch, "/backend-api/codex/alpha/search", service.PlatformOpenAI, EndpointAlphaSearch},
{"openai image generations", EndpointImagesGenerations, "/v1/images/generations", service.PlatformOpenAI, EndpointImagesGenerations},
{"openai image edits", EndpointImagesEdits, "/openai/v1/images/edits", service.PlatformOpenAI, EndpointImagesEdits},
{"grok chat defaults to responses without runtime result", EndpointChatCompletions, "/v1/chat/completions", service.PlatformGrok, EndpointResponses},
{"grok responses", EndpointResponses, "/v1/responses", service.PlatformGrok, EndpointResponses},
{"grok video generations", EndpointVideosGenerations, "/v1/videos/generations", service.PlatformGrok, EndpointVideosGenerations},
{"grok video status", EndpointVideos, "/videos/req_123", service.PlatformGrok, EndpointVideos},
@@ -138,6 +144,59 @@ func TestDeriveUpstreamEndpoint(t *testing.T) {
}
}
func TestResolveOpenAIUpstreamEndpointPrefersForwardResult(t *testing.T) {
tests := []struct {
name string
account *service.Account
result *service.OpenAIForwardResult
runtimeEndpoint string
want string
}{
{
name: "grok raw chat result overrides stale context",
account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth},
result: &service.OpenAIForwardResult{UpstreamEndpoint: EndpointChatCompletions},
runtimeEndpoint: EndpointResponses,
want: EndpointChatCompletions,
},
{
name: "grok chat bridged to responses",
account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth},
result: &service.OpenAIForwardResult{UpstreamEndpoint: EndpointResponses},
want: EndpointResponses,
},
{
name: "grok empty result keeps responses default",
account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth},
result: &service.OpenAIForwardResult{},
want: EndpointResponses,
},
{
name: "grok raw error uses runtime endpoint",
account: &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeOAuth},
runtimeEndpoint: EndpointChatCompletions,
want: EndpointChatCompletions,
},
{
name: "openai behavior remains responses",
account: &service.Account{Platform: service.PlatformOpenAI, Type: service.AccountTypeOAuth},
result: &service.OpenAIForwardResult{},
want: EndpointResponses,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, EndpointChatCompletions, nil)
c.Set(ctxKeyInboundEndpoint, EndpointChatCompletions)
service.SetActualOpenAIUpstreamEndpoint(c, tt.runtimeEndpoint)
require.Equal(t, tt.want, resolveOpenAIUpstreamEndpoint(c, tt.account, tt.result))
})
}
}
// ──────────────────────────────────────────────────────────
// responsesSubpathSuffix
// ──────────────────────────────────────────────────────────
@@ -107,3 +107,31 @@ func classifyNoAccountErrorFromGin(
}
return classifyNoAccountError(ctx, diag, apiKey, routingModel, displayModel, platform)
}
func classifyOpenAICompatibleNoAccountErrorFromGin(
c *gin.Context,
diag service.ModelAvailabilityDiagnoser,
apiKey *service.APIKey,
routingModel string,
displayModel string,
) noAccountErrorClassification {
return classifyNoAccountErrorFromGin(
c,
diag,
apiKey,
routingModel,
displayModel,
openAICompatibleRequestPlatform(apiKey),
)
}
func openAICompatibleSelectionErrorForLog(err error, platform string) error {
if err == nil || platform != service.PlatformGrok {
return err
}
message := strings.ReplaceAll(err.Error(), "OpenAI accounts", "Grok accounts")
if message == err.Error() {
return err
}
return fmt.Errorf("%s", message)
}
@@ -4,6 +4,7 @@ package handler
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"testing"
@@ -114,6 +115,33 @@ func TestClassifyNoAccountError_ModelNotSupported_Returns404(t *testing.T) {
require.Equal(t, int64(42), *fd.calls[0].GroupID)
}
func TestClassifyOpenAICompatibleNoAccountError_GrokUsesGrokPlatform(t *testing.T) {
c := newTestGinContextWithRequest()
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: false}}
groupID := int64(43)
apiKey := &service.APIKey{
GroupID: &groupID,
Group: &service.Group{
ID: groupID,
Platform: service.PlatformGrok,
},
}
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, fd, apiKey, "grok-4.5", "grok-4.5")
require.Equal(t, http.StatusNotFound, cls.Status)
require.Equal(t, "model_not_found", cls.ErrType)
require.True(t, cls.ModelNotFound)
require.Len(t, fd.calls, 1)
require.Equal(t, service.PlatformGrok, fd.calls[0].Platform)
logErr := openAICompatibleSelectionErrorForLog(
fmt.Errorf("no available OpenAI accounts supporting model: grok-4.5"),
service.PlatformGrok,
)
require.EqualError(t, logErr, "no available Grok accounts supporting model: grok-4.5")
}
func TestClassifyNoAccountError_HasModelSupport_KeepsRoutingMessageGenerationToCaller(t *testing.T) {
c := newTestGinContextWithRequest()
fd := &fakeDiagnoser{resp: service.ModelAvailabilityDiagnosis{HasAccountsInPool: true, HasModelSupport: true}}
@@ -0,0 +1,250 @@
package handler
import (
"context"
"errors"
"net/http"
"strconv"
"strings"
"time"
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
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"
)
// AlphaSearch proxies the standalone search endpoint used by Codex Responses Lite.
func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
streamStarted := false
defer h.recoverResponsesPanic(c, &streamStarted)
setOpenAIClientTransportHTTP(c)
requestStart := time.Now()
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok || apiKey.Group == nil {
h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
return
}
if apiKey.Group.Platform != service.PlatformOpenAI {
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex alpha search is only available for OpenAI groups")
return
}
subject, ok := middleware2.GetAuthSubjectFromContext(c)
if !ok {
h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found")
return
}
reqLog := requestLogger(
c,
"handler.openai_gateway.alpha_search",
zap.Int64("user_id", subject.UserID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
)
if !h.ensureResponsesDependencies(c, reqLog) {
return
}
body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
if err != nil {
if maxErr, ok := extractMaxBytesError(err); ok {
h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
return
}
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
return
}
if len(body) == 0 {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
return
}
if !gjson.ValidBytes(body) {
logRequestBodyParseFailure(reqLog, body, nil)
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
}
modelResult := gjson.GetBytes(body, "model")
if !modelResult.Exists() || modelResult.Type != gjson.String || strings.TrimSpace(modelResult.String()) == "" {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
return
}
requestedModel := strings.TrimSpace(modelResult.String())
reqLog = reqLog.With(zap.String("model", requestedModel))
setOpsRequestContext(c, requestedModel, false)
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestedModel)
forwardBody := openAIModelMappedBody(body, channelMapping.Mapped, channelMapping.MappedModel, h.gatewayService.ReplaceModelInBody)
subscription, _ := middleware2.GetSubscriptionFromContext(c)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
userRelease, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog)
if !acquired {
return
}
if userRelease != nil {
defer userRelease()
}
if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
status, code, message, retryAfter := billingErrorDetails(err)
if retryAfter > 0 {
c.Header("Retry-After", strconv.Itoa(retryAfter))
}
h.errorResponse(c, status, code, message)
return
}
searchID := strings.TrimSpace(gjson.GetBytes(body, "id").String())
sessionHash := h.gatewayService.GenerateSessionHashWithFallback(c, nil, searchID)
failedAccountIDs := make(map[int64]struct{})
var lastFailoverErr *service.UpstreamFailoverError
switchCount := 0
routingStart := time.Now()
for {
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
"",
sessionHash,
requestedModel,
failedAccountIDs,
service.OpenAIUpstreamTransportHTTPSSE,
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
service.PlatformOpenAI,
)
if err != nil || selection == nil || selection.Account == nil {
if len(failedAccountIDs) == 0 {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestedModel, requestedModel, service.PlatformOpenAI)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
return
}
if lastFailoverErr != nil {
h.handleFailoverExhausted(c, lastFailoverErr, false)
} else {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
return
}
account := selection.Account
setOpsSelectedAccount(c, account.ID, account.Platform)
accountRelease, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog)
if !acquired {
return
}
service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds())
writerSizeBeforeForward := c.Writer.Size()
forwardStart := time.Now()
var result *service.OpenAIForwardResult
result, err = func() (*service.OpenAIForwardResult, error) {
if accountRelease != nil {
defer accountRelease()
}
return h.gatewayService.ForwardAlphaSearch(c.Request.Context(), c, account, forwardBody)
}()
service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds())
if err == nil {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil)
if result != nil {
h.recordAlphaSearchUsage(c, apiKey, account, subscription, channelMapping, requestedModel, body, result, subject.UserID)
}
return
}
var failoverErr *service.UpstreamFailoverError
if !errors.As(err, &failoverErr) {
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
if c.Writer.Size() == writerSizeBeforeForward {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
}
reqLog.Warn("openai_alpha_search.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err))
return
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
if c.Writer.Size() != writerSizeBeforeForward {
h.handleFailoverExhausted(c, failoverErr, true)
return
}
h.gatewayService.RecordOpenAIAccountSwitch()
failedAccountIDs[account.ID] = struct{}{}
lastFailoverErr = failoverErr
if switchCount >= h.maxAccountSwitches {
h.handleFailoverExhausted(c, failoverErr, false)
return
}
switchCount++
if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) {
h.handleFailoverExhausted(c, failoverErr, false)
return
}
reqLog.Warn("openai_alpha_search.upstream_failover_switching",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("switch_count", switchCount),
)
}
}
// recordAlphaSearchUsage 为一次成功的 alpha/search 网页搜索落按次计费用量行
// (上游不返回 usage 字段,按 WebSearchCalls 走分组单价 × 倍率的按次口径)。
// 与 images 一致使用 mandatory 池提交,池满时同步兜底执行,保证扣费不丢。
func (h *OpenAIGatewayHandler) recordAlphaSearchUsage(
c *gin.Context,
apiKey *service.APIKey,
account *service.Account,
subscription *service.UserSubscription,
channelMapping service.ChannelMappingResult,
requestedModel string,
body []byte,
result *service.OpenAIForwardResult,
userID int64,
) {
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
requestPayloadHash := service.HashUsageRequestPayload(body)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
Result: result,
APIKey: apiKey,
User: apiKey.User,
Account: account,
Subscription: subscription,
InboundEndpoint: inboundEndpoint,
UpstreamEndpoint: upstreamEndpoint,
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
ChannelUsageFields: channelMapping.ToUsageFields(requestedModel, result.UpstreamModel),
}); err != nil {
logger.L().With(
zap.String("component", "handler.openai_gateway.alpha_search"),
zap.Int64("user_id", userID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
zap.String("model", requestedModel),
zap.Int64("account_id", account.ID),
).Error("openai_alpha_search.record_usage_failed", zap.Error(err))
}
})
}
@@ -5,6 +5,7 @@ import (
"errors"
"net/http"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
@@ -150,11 +151,11 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
)
if err != nil {
reqLog.Warn("openai_chat_completions.account_select_failed",
zap.Error(err),
zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)),
zap.Int("excluded_account_count", len(failedAccountIDs)),
)
if len(failedAccountIDs) == 0 {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI)
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -170,7 +171,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}
}
if selection == nil || selection.Account == nil {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI)
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
}
@@ -298,7 +299,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
@@ -337,14 +338,22 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}
// resolveOpenAIUpstreamEndpoint returns the actual upstream endpoint for an
// OpenAI account, used by every OpenAI usage-recording site. APIKey accounts
// whose upstream is forced or probed to not support the Responses API are
// served directly via /v1/chat/completions (the raw chat path) regardless of
// the inbound endpoint; everything else goes through the Responses API.
func resolveOpenAIUpstreamEndpoint(c *gin.Context, account *service.Account) string {
// OpenAI-compatible account. A forwarding result is authoritative because a
// single inbound route may choose raw Chat or a Responses bridge at runtime.
// The account-based derivation remains as a fallback for existing callers and
// forwarding paths that do not report their endpoint yet.
func resolveOpenAIUpstreamEndpoint(c *gin.Context, account *service.Account, result *service.OpenAIForwardResult) string {
if result != nil {
if endpoint := strings.TrimSpace(result.UpstreamEndpoint); endpoint != "" {
return endpoint
}
}
if endpoint := service.GetActualOpenAIUpstreamEndpoint(c); endpoint != "" {
return endpoint
}
if account != nil && account.Type == service.AccountTypeAPIKey &&
!openai_compat.ShouldUseResponsesAPI(account.Extra) {
return "/v1/chat/completions"
return EndpointChatCompletions
}
return GetUpstreamEndpoint(c, account.Platform)
}
@@ -23,46 +23,61 @@ func newCompactBodySignalTestContext(t *testing.T, path string, body []byte) *gi
return c
}
// body-signal 提升后必须与 path-based compact 走同一条链路:
// path 改写、requireCompact 判定、stream/store/prompt_cache_key 归一化删除。
// 回归防护:若 stream 字段存活,Forward 会用流式 handler 解析 compact 的
// JSON 响应,导致 "stream ended before a terminal event" 的换号 failover 风暴。
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalPromoted(t *testing.T) {
func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2StaysOnResponses(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{
"model":"gpt-5.5",
"model":"gpt-5.6-sol",
"stream":true,
"store":true,
"prompt_cache_key":"pck-signal-1",
"reasoning":{"effort":"max","context":"all_turns"},
"input":[
{"type":"message","role":"user","content":"hello"},
{"type":"compaction_trigger"}
]
}`)
c := newCompactBodySignalTestContext(t, "/v1/responses", body)
c.Request.Header.Set("x-codex-beta-features", "responses_websockets_v2, remote_compaction_v2, another_feature")
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
require.Equal(t, "/v1/responses/compact", c.Request.URL.Path)
require.True(t, isOpenAIRemoteCompactPath(c))
require.False(t, gjson.GetBytes(normalized, "stream").Exists())
require.False(t, gjson.GetBytes(normalized, "store").Exists())
require.False(t, gjson.GetBytes(normalized, "prompt_cache_key").Exists())
require.Equal(t, "gpt-5.5", gjson.GetBytes(normalized, "model").String())
require.True(t, gjson.GetBytes(normalized, "input").IsArray())
require.Equal(t, "/v1/responses", c.Request.URL.Path)
require.False(t, isOpenAIRemoteCompactPath(c))
require.Equal(t, body, normalized)
require.True(t, gjson.GetBytes(normalized, "stream").Bool())
require.True(t, gjson.GetBytes(normalized, "store").Bool())
require.Equal(t, "pck-signal-1", gjson.GetBytes(normalized, "prompt_cache_key").String())
require.Equal(t, "max", gjson.GetBytes(normalized, "reasoning.effort").String())
require.Equal(t, "all_turns", gjson.GetBytes(normalized, "reasoning.context").String())
reqStream, streamOK := parseOpenAICompatibleStream(normalized)
require.True(t, streamOK)
require.False(t, reqStream)
require.True(t, reqStream)
seed, exists := c.Get(service.OpenAICompactSessionSeedKeyForTest())
require.True(t, exists)
require.Equal(t, "pck-signal-1", seed)
_, seedExists := c.Get(service.OpenAICompactSessionSeedKeyForTest())
require.False(t, seedExists)
_, streamMarkerExists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.False(t, streamMarkerExists)
}
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlash(t *testing.T) {
func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2PathAliasesStayOnResponses(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.6-sol","stream":true,"input":[{"type":"compaction_trigger"}]}`)
for _, path := range []string{"/v1/responses/", "/backend-api/codex/responses"} {
t.Run(path, func(t *testing.T) {
c := newCompactBodySignalTestContext(t, path, body)
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
require.Equal(t, path, c.Request.URL.Path)
require.Equal(t, body, normalized)
})
}
}
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlashPromoted(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`)
c := newCompactBodySignalTestContext(t, "/v1/responses/", body)
@@ -82,6 +97,64 @@ func TestNormalizeOpenAIResponsesCompactRequest_CodexDirectAliasPromoted(t *test
require.Equal(t, "/backend-api/codex/responses/compact", c.Request.URL.Path)
}
func TestNormalizeOpenAIResponsesCompactRequest_NonRemoteV2BodySignalPromoted(t *testing.T) {
h := &OpenAIGatewayHandler{}
tests := []struct {
name string
body []byte
betaHeader string
wantMarked bool
}{
{
name: "no_header",
body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`),
wantMarked: true,
},
{
name: "unrelated_header",
body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`),
betaHeader: "responses_websockets_v2",
wantMarked: true,
},
{
name: "wrong_case_header",
body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`),
betaHeader: "REMOTE_COMPACTION_V2",
wantMarked: true,
},
{
name: "stream_false",
body: []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`),
betaHeader: "remote_compaction_v2",
},
{
name: "stream_absent",
body: []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`),
betaHeader: "remote_compaction_v2",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
c := newCompactBodySignalTestContext(t, "/v1/responses", tt.body)
if tt.betaHeader != "" {
c.Request.Header.Set("x-codex-beta-features", tt.betaHeader)
}
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), tt.body)
require.True(t, ok)
require.Equal(t, "/v1/responses/compact", c.Request.URL.Path)
require.False(t, gjson.GetBytes(normalized, "stream").Exists())
marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.Equal(t, tt.wantMarked, exists)
if tt.wantMarked {
require.Equal(t, true, marked)
}
})
}
}
func TestNormalizeOpenAIResponsesCompactRequest_NoTriggerUntouched(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"message","role":"user","content":"hello"}]}`)
@@ -99,6 +172,7 @@ func TestNormalizeOpenAIResponsesCompactRequest_PathBasedNoDoubleSuffix(t *testi
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","stream":true,"store":true,"input":[{"type":"message","role":"user","content":"hello"}]}`)
c := newCompactBodySignalTestContext(t, "/v1/responses/compact", body)
c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2")
normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
@@ -118,36 +192,6 @@ func TestNormalizeOpenAIResponsesCompactRequest_SubpathNotPromoted(t *testing.T)
require.Equal(t, body, normalized)
}
// 回归 #3875:body-signal 原始请求 stream:true 时必须标记 client-stream,
// 供响应写回阶段把上游 unary JSON 合成回 Codex remote compact v2 所需的 SSE。
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamTrueMarksClientStream(t *testing.T) {
h := &OpenAIGatewayHandler{}
body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`)
c := newCompactBodySignalTestContext(t, "/v1/responses", body)
_, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok)
marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.True(t, exists)
require.Equal(t, true, marked)
}
func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamFalseNotMarked(t *testing.T) {
h := &OpenAIGatewayHandler{}
for name, body := range map[string][]byte{
"stream_false": []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`),
"stream_absent": []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`),
} {
c := newCompactBodySignalTestContext(t, "/v1/responses", body)
_, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body)
require.True(t, ok, name)
require.Equal(t, "/v1/responses/compact", c.Request.URL.Path, name)
_, exists := c.Get(service.OpenAICompactClientStreamKeyForTest())
require.False(t, exists, "case %s 不应标记 client-stream", name)
}
}
// path-based compact(Codex v1 unary 协议)即使 body 带 stream:true 也不标记,
// 保持 JSON 写回行为不变。
func TestNormalizeOpenAIResponsesCompactRequest_PathBasedStreamTrueNotMarked(t *testing.T) {
@@ -115,8 +115,9 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
if err != nil {
reqLog.Warn("openai_count_tokens.account_select_failed", zap.Error(err))
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI)
requestPlatform := openAICompatibleRequestPlatform(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 {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -124,7 +125,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
return
}
if selection == nil || selection.Account == nil {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI)
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
}
@@ -351,7 +351,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
)
if err != nil {
reqLog.Warn("openai.account_select_failed",
zap.Error(err),
zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)),
zap.Int("excluded_account_count", len(failedAccountIDs)),
)
if len(failedAccountIDs) == 0 {
@@ -360,7 +360,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "compact_not_supported", "No available OpenAI accounts support /responses/compact", streamStarted)
return
}
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI)
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -375,7 +375,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
return
}
if selection == nil || selection.Account == nil {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel, service.PlatformOpenAI)
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, reqModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
}
@@ -522,7 +522,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
clientIP := ip.GetClientIP(c)
requestPayloadHash := service.HashUsageRequestPayload(body)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
@@ -580,21 +580,33 @@ func isBareOpenAIResponsesPath(c *gin.Context) bool {
return strings.HasSuffix(normalizedPath, "/responses")
}
// normalizeOpenAIResponsesCompactRequest 统一处理两种入站 compact 形态:
// path-based(POST /v1/responses/compact)与 Codex remote compact v2 的
// body-signal(普通 POST /v1/responses 的 input 中携带 type=compaction_trigger,
// 见 #3777)。body-signal 命中时在 stream 解析、compact body 归一化与
// requireCompact 调度判定之前改写 URL path,使后续全部链路(含 passthrough
// 分支与上游 URL 构建)与 path-based 完全一致。
func isOpenAIRemoteCompactionV2Request(c *gin.Context, body []byte) bool {
stream, valid := parseOpenAICompatibleStream(body)
if !valid || !stream || c == nil || c.Request == nil {
return false
}
for _, header := range c.Request.Header.Values("x-codex-beta-features") {
for _, feature := range strings.Split(header, ",") {
if strings.TrimSpace(feature) == "remote_compaction_v2" {
return true
}
}
}
return false
}
// normalizeOpenAIResponsesCompactRequest keeps Codex remote compaction v2 on
// its native streaming /responses wire and preserves the legacy body-signal
// promotion for clients that do not explicitly advertise that protocol.
// 返回归一化后的 body;ok=false 表示错误响应已写出,调用方应直接 return。
func (h *OpenAIGatewayHandler) normalizeOpenAIResponsesCompactRequest(c *gin.Context, reqLog *zap.Logger, body []byte) ([]byte, bool) {
isCompactRequest := service.IsOpenAIResponsesCompactPathForTest(c)
if !isCompactRequest && isBareOpenAIResponsesPath(c) && service.HasCompactionTriggerInInput(body) {
if isOpenAIRemoteCompactionV2Request(c, body) {
return body, true
}
c.Request.URL.Path = strings.TrimRight(c.Request.URL.Path, "/") + "/compact"
isCompactRequest = true
// Codex remote compact v2 的原始请求是流式 /responses:白名单归一化会删除
// stream 并让上游走 unary JSON,但客户端仍按 SSE 消费响应。记录原始
// stream 意图,响应写回阶段据此把 JSON 合成回 SSE(#3875)。
clientStream := gjson.GetBytes(body, "stream").Bool()
if clientStream {
service.MarkOpenAICompactClientStream(c)
@@ -843,12 +855,12 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
)
if err != nil {
reqLog.Warn("openai_messages.account_select_failed",
zap.Error(err),
zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)),
zap.Int("excluded_account_count", len(failedAccountIDs)),
)
if len(failedAccountIDs) == 0 {
if err != nil {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI)
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
@@ -865,7 +877,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
}
}
if selection == nil || selection.Account == nil {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI)
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
}
@@ -994,7 +1006,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
clientIP := ip.GetClientIP(c)
requestPayloadHash := service.HashUsageRequestPayload(body)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
@@ -1444,7 +1456,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
)
if err != nil {
reqLog.Warn("openai.websocket_account_select_failed",
zap.Error(err),
zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)),
zap.Int("excluded_account_count", len(failedAccountIDs)),
)
if lastFailoverErr != nil {
@@ -1601,7 +1613,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
}
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, result.FirstTokenMs)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) {
@@ -2441,7 +2453,7 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey
var accountID int64
if account != nil {
accountID = account.ID
upstreamEndpoint = resolveOpenAIUpstreamEndpoint(c, account)
upstreamEndpoint = resolveOpenAIUpstreamEndpoint(c, account, nil)
}
stream := false
if v, ok := c.Get(opsStreamKey); ok {
@@ -1,9 +1,15 @@
package handler
import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
@@ -57,3 +63,28 @@ func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) {
assert.Nil(t, p)
})
}
func TestOpsCaptureWriter_CompactKeepaliveRestoresOriginalWriter(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
outerStatus := -1
router.Use(func(c *gin.Context) {
c.Next()
outerStatus = c.Writer.Status()
})
router.Use(OpsErrorLoggerMiddleware(nil))
router.GET("/compact", func(c *gin.Context) {
service.MarkOpenAICompactClientStream(c)
stop := service.StartOpenAICompactSSEKeepalive(c, time.Hour)
defer stop()
c.Status(http.StatusOK)
})
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "/compact", nil)
require.NotPanics(t, func() {
router.ServeHTTP(recorder, request)
})
require.Equal(t, http.StatusOK, outerStatus)
require.Equal(t, http.StatusOK, recorder.Code)
}