mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
Merge pull request #5810 from Pluviobyte/codex/fix-responses-input-tokens
fix(codex): handle Responses input token preflight
This commit is contained in:
@@ -15,20 +15,21 @@ import (
|
||||
// ──────────────────────────────────────────────────────────
|
||||
|
||||
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"
|
||||
EndpointImagesEdits = "/v1/images/edits"
|
||||
EndpointImageTasks = "/v1/images/tasks"
|
||||
EndpointVideosGenerations = "/v1/videos/generations"
|
||||
EndpointVideosEdits = "/v1/videos/edits"
|
||||
EndpointVideosExtensions = "/v1/videos/extensions"
|
||||
EndpointVideos = "/v1/videos"
|
||||
EndpointGeminiModels = "/v1beta/models"
|
||||
EndpointMessages = "/v1/messages"
|
||||
EndpointChatCompletions = "/v1/chat/completions"
|
||||
EndpointEmbeddings = "/v1/embeddings"
|
||||
EndpointAlphaSearch = "/v1/alpha/search"
|
||||
EndpointResponses = "/v1/responses"
|
||||
EndpointResponsesCompact = "/v1/responses/compact"
|
||||
EndpointResponsesInputTokens = "/v1/responses/input_tokens"
|
||||
EndpointImagesGenerations = "/v1/images/generations"
|
||||
EndpointImagesEdits = "/v1/images/edits"
|
||||
EndpointImageTasks = "/v1/images/tasks"
|
||||
EndpointVideosGenerations = "/v1/videos/generations"
|
||||
EndpointVideosEdits = "/v1/videos/edits"
|
||||
EndpointVideosExtensions = "/v1/videos/extensions"
|
||||
EndpointVideos = "/v1/videos"
|
||||
EndpointGeminiModels = "/v1beta/models"
|
||||
)
|
||||
|
||||
const EndpointAntigravityGenerateContent = "/v1internal:streamGenerateContent"
|
||||
@@ -80,6 +81,8 @@ const (
|
||||
func NormalizeInboundEndpoint(path string) string {
|
||||
path = strings.TrimSpace(path)
|
||||
switch {
|
||||
case strings.Contains(path, EndpointResponsesInputTokens) || isResponsesInputTokensAliasPath(path):
|
||||
return EndpointResponsesInputTokens
|
||||
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"):
|
||||
@@ -113,6 +116,15 @@ func NormalizeInboundEndpoint(path string) string {
|
||||
}
|
||||
}
|
||||
|
||||
func isResponsesInputTokensAliasPath(path string) bool {
|
||||
trimmed := strings.TrimRight(strings.TrimSpace(path), "/")
|
||||
if trimmed == "" {
|
||||
return false
|
||||
}
|
||||
return isBareOrSubpathOf(trimmed, "/responses/input_tokens") ||
|
||||
isBareOrSubpathOf(trimmed, "/backend-api/codex/responses/input_tokens")
|
||||
}
|
||||
|
||||
// isResponsesCompactAliasPath reports whether path is the bare/alias
|
||||
// "compact" client endpoint — i.e. it is rooted at "/responses/compact"
|
||||
// or "/backend-api/codex/responses/compact" (bare routes that serve
|
||||
@@ -185,7 +197,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
|
||||
|
||||
switch platform {
|
||||
case service.PlatformOpenAI, service.PlatformGrok:
|
||||
if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideosEdits || inbound == EndpointVideosExtensions || inbound == EndpointVideos {
|
||||
if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointResponsesInputTokens || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideosEdits || inbound == EndpointVideosExtensions || inbound == EndpointVideos {
|
||||
return inbound
|
||||
}
|
||||
// OpenAI forwards everything to the Responses API.
|
||||
|
||||
@@ -27,6 +27,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
|
||||
{"/v1/embeddings", EndpointEmbeddings},
|
||||
{"/v1/alpha/search", EndpointAlphaSearch},
|
||||
{"/v1/responses", EndpointResponses},
|
||||
{"/v1/responses/input_tokens", EndpointResponsesInputTokens},
|
||||
{"/v1/responses/compact", EndpointResponsesCompact},
|
||||
{"/v1/responses/compact/detail", EndpointResponsesCompact},
|
||||
{"/v1/images/generations", EndpointImagesGenerations},
|
||||
@@ -50,6 +51,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
|
||||
|
||||
// Bare top-level alias route "/responses" — root vs. compact.
|
||||
{"/responses", EndpointResponses},
|
||||
{"/responses/input_tokens", EndpointResponsesInputTokens},
|
||||
{"/responses/compact", EndpointResponsesCompact},
|
||||
{"/responses/compact/detail", EndpointResponsesCompact},
|
||||
{"/alpha/search", EndpointAlphaSearch},
|
||||
@@ -57,6 +59,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
|
||||
|
||||
// Bare Codex direct alias route — root vs. compact.
|
||||
{"/backend-api/codex/responses", EndpointResponses},
|
||||
{"/backend-api/codex/responses/input_tokens", EndpointResponsesInputTokens},
|
||||
{"/backend-api/codex/responses/compact", EndpointResponsesCompact},
|
||||
{"/backend-api/codex/responses/compact/detail", EndpointResponsesCompact},
|
||||
{"/backend-api/codex/alpha/search", EndpointAlphaSearch},
|
||||
@@ -100,6 +103,7 @@ func TestDeriveUpstreamEndpoint(t *testing.T) {
|
||||
|
||||
// OpenAI — root Responses.
|
||||
{"openai responses root", EndpointResponses, "/v1/responses", service.PlatformOpenAI, EndpointResponses},
|
||||
{"openai responses input tokens", EndpointResponsesInputTokens, "/v1/responses/input_tokens", service.PlatformOpenAI, EndpointResponsesInputTokens},
|
||||
|
||||
// OpenAI — compact, raw path carries the derivable "/compact"
|
||||
// (or nested) suffix, which must be preserved on the upstream
|
||||
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/domain"
|
||||
@@ -10,9 +11,134 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/tidwall/gjson"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
// ResponsesInputTokens handles native OpenAI POST
|
||||
// /v1/responses/input_tokens requests without routing them through the normal
|
||||
// Responses generation and usage-recording pipeline.
|
||||
func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) {
|
||||
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
|
||||
if !ok {
|
||||
h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
|
||||
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.responses_input_tokens",
|
||||
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 := readLenientJSONRequestBodyWithPrealloc(c.Request, h.cfg)
|
||||
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 || !gjson.ValidBytes(body) {
|
||||
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
|
||||
}
|
||||
reqModel := strings.TrimSpace(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
|
||||
}
|
||||
|
||||
setOpsRequestContext(c, reqModel, false)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(false, false)))
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
subscription, _ := middleware2.GetSubscriptionFromContext(c)
|
||||
if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
|
||||
reqLog.Info("openai_input_tokens.billing_eligibility_check_failed", zap.Error(err))
|
||||
status, code, message, retryAfter := billingErrorDetails(err)
|
||||
if retryAfter > 0 {
|
||||
c.Header("Retry-After", strconv.Itoa(retryAfter))
|
||||
}
|
||||
h.errorResponse(c, status, code, message)
|
||||
return
|
||||
}
|
||||
|
||||
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel)
|
||||
routingModel := reqModel
|
||||
forwardBody := body
|
||||
if channelMapping.Mapped {
|
||||
routingModel = channelMapping.MappedModel
|
||||
forwardBody = h.gatewayService.ReplaceModelInBody(body, routingModel)
|
||||
}
|
||||
|
||||
// Token counting is not billed, so it must not be excluded by the profit gate.
|
||||
c.Request = c.Request.WithContext(service.WithOpenAIProfitControlSuppressed(c.Request.Context()))
|
||||
requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey)
|
||||
sessionHash := h.gatewayService.GenerateSessionHash(c, body)
|
||||
requestStart := time.Now()
|
||||
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
|
||||
c.Request.Context(),
|
||||
apiKey.GroupID,
|
||||
"",
|
||||
sessionHash,
|
||||
routingModel,
|
||||
nil,
|
||||
service.OpenAIUpstreamTransportAny,
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
requestPlatform,
|
||||
)
|
||||
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
|
||||
if err != nil {
|
||||
reqLog.Warn("openai_input_tokens.account_select_failed", zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)))
|
||||
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, routingModel, reqModel)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
|
||||
}
|
||||
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
|
||||
return
|
||||
}
|
||||
if selection == nil || selection.Account == nil {
|
||||
cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, routingModel, reqModel)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
}
|
||||
h.errorResponse(c, cls.Status, cls.ErrType, cls.Message)
|
||||
return
|
||||
}
|
||||
|
||||
account := selection.Account
|
||||
setOpsSelectedAccount(c, account.ID, account.Platform)
|
||||
if selection.Acquired && selection.ReleaseFunc != nil {
|
||||
defer selection.ReleaseFunc()
|
||||
}
|
||||
if err := h.gatewayService.ForwardResponsesInputTokens(c.Request.Context(), c, account, forwardBody); err != nil {
|
||||
reqLog.Error("openai_input_tokens.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
// GrokCountTokens handles Anthropic-compatible count_tokens requests locally.
|
||||
// The route middleware already authenticates the API key and resolves the
|
||||
// group; this handler intentionally does not select an account or check billing.
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestResponsesInputTokensIsClassifiedAsTokenCountRequest(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/input_tokens", nil)
|
||||
|
||||
require.True(t, isCountTokensRequest(c))
|
||||
require.True(t, isTokenCountRequestPath("/responses/input_tokens"))
|
||||
require.False(t, isTokenCountRequestPath("/v1/responses"))
|
||||
}
|
||||
@@ -1249,7 +1249,11 @@ func isCountTokensRequest(c *gin.Context) bool {
|
||||
if c == nil || c.Request == nil || c.Request.URL == nil {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(c.Request.URL.Path, "/count_tokens")
|
||||
return isTokenCountRequestPath(c.Request.URL.Path)
|
||||
}
|
||||
|
||||
func isTokenCountRequestPath(path string) bool {
|
||||
return strings.Contains(path, "/count_tokens") || strings.Contains(path, "/responses/input_tokens")
|
||||
}
|
||||
|
||||
func applyOpsLatencyFieldsFromContext(c *gin.Context, entry *service.OpsInsertErrorLogInput) {
|
||||
@@ -1774,7 +1778,7 @@ func shouldSkipOpsErrorLog(ctx context.Context, ops *service.OpsService, message
|
||||
bodyLower := strings.ToLower(body)
|
||||
|
||||
// Check if count_tokens errors should be ignored
|
||||
if settings.IgnoreCountTokensErrors && strings.Contains(requestPath, "/count_tokens") {
|
||||
if settings.IgnoreCountTokensErrors && isTokenCountRequestPath(requestPath) {
|
||||
return true
|
||||
}
|
||||
|
||||
|
||||
@@ -170,6 +170,10 @@ func RegisterGatewayRoutes(
|
||||
})
|
||||
return
|
||||
}
|
||||
if service.IsOpenAIResponsesInputTokensRequestPath(c) && isOpenAIResponsesCompatibleGatewayPlatform(c) {
|
||||
h.OpenAIGateway.ResponsesInputTokens(c)
|
||||
return
|
||||
}
|
||||
next(c)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
|
||||
@@ -39,6 +40,177 @@ type openAIInputTokensCountPrepared struct {
|
||||
UpstreamModel string
|
||||
}
|
||||
|
||||
// ForwardResponsesInputTokens handles the native OpenAI
|
||||
// POST /v1/responses/input_tokens shape. Custom OpenAI-compatible relays often
|
||||
// implement /responses but not this preflight endpoint, so those accounts use
|
||||
// the local estimator instead of receiving a request that is known to fail.
|
||||
func (s *OpenAIGatewayService) ForwardResponsesInputTokens(
|
||||
ctx context.Context,
|
||||
c *gin.Context,
|
||||
account *Account,
|
||||
body []byte,
|
||||
) error {
|
||||
if account == nil {
|
||||
writeOpenAIResponsesInputTokensError(c, http.StatusServiceUnavailable, "api_error", "No available OpenAI accounts")
|
||||
return fmt.Errorf("responses input_tokens: missing account")
|
||||
}
|
||||
|
||||
prepared, err := prepareNativeOpenAIInputTokensCountRequest(body, account)
|
||||
if err != nil {
|
||||
writeOpenAIResponsesInputTokensError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
|
||||
return err
|
||||
}
|
||||
|
||||
if shouldEstimateOpenAIInputTokensLocally(account) {
|
||||
writeOpenAIResponsesInputTokensFallback(c, account, prepared, 0, "custom_relay")
|
||||
return nil
|
||||
}
|
||||
|
||||
token, _, err := s.GetAccessToken(ctx, account)
|
||||
if err != nil {
|
||||
writeOpenAIResponsesInputTokensError(c, http.StatusBadGateway, "upstream_error", "Failed to get access token")
|
||||
return fmt.Errorf("responses input_tokens: get access token: %w", err)
|
||||
}
|
||||
|
||||
upstreamBody := ReplaceModelInBody(body, prepared.UpstreamModel)
|
||||
upstreamReq, err := s.buildInputTokensUpstreamRequest(ctx, c, account, upstreamBody, token)
|
||||
if err != nil {
|
||||
writeOpenAIResponsesInputTokensError(c, http.StatusInternalServerError, "api_error", "Failed to build request")
|
||||
return fmt.Errorf("responses input_tokens: build upstream request: %w", err)
|
||||
}
|
||||
|
||||
proxyURL := ""
|
||||
if account.Proxy != nil {
|
||||
proxyURL = account.Proxy.URL()
|
||||
}
|
||||
resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency)
|
||||
if err != nil {
|
||||
safeErr := sanitizeUpstreamErrorMessage(err.Error())
|
||||
setOpsUpstreamError(c, 0, safeErr, "")
|
||||
writeOpenAIResponsesInputTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
|
||||
return fmt.Errorf("responses input_tokens: upstream request failed: %s", safeErr)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
writeOpenAIResponsesInputTokensError(c, http.StatusBadGateway, "upstream_error", "Failed to read response")
|
||||
return fmt.Errorf("responses input_tokens: read upstream response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode >= 400 {
|
||||
if isOpenAIResponsesInputTokensUnsupported(account, resp.StatusCode, respBody) {
|
||||
writeOpenAIResponsesInputTokensFallback(c, account, prepared, resp.StatusCode, "upstream_unsupported")
|
||||
return nil
|
||||
}
|
||||
if s.rateLimitService != nil {
|
||||
s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
}
|
||||
upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
|
||||
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, "")
|
||||
writeOpenAIResponsesInputTokensError(c, resp.StatusCode, "upstream_error", "Upstream request failed")
|
||||
if upstreamMsg == "" {
|
||||
return fmt.Errorf("responses input_tokens: upstream error: %d", resp.StatusCode)
|
||||
}
|
||||
return fmt.Errorf("responses input_tokens: upstream error: %d message=%s", resp.StatusCode, upstreamMsg)
|
||||
}
|
||||
|
||||
inputTokens := gjson.GetBytes(respBody, "input_tokens")
|
||||
if !inputTokens.Exists() {
|
||||
writeOpenAIResponsesInputTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream response missing input_tokens")
|
||||
return fmt.Errorf("responses input_tokens: upstream response missing input_tokens")
|
||||
}
|
||||
contentType := strings.TrimSpace(resp.Header.Get("Content-Type"))
|
||||
if contentType == "" {
|
||||
contentType = "application/json"
|
||||
}
|
||||
c.Data(http.StatusOK, contentType, respBody)
|
||||
return nil
|
||||
}
|
||||
|
||||
func prepareNativeOpenAIInputTokensCountRequest(body []byte, account *Account) (*openAIInputTokensCountPrepared, error) {
|
||||
var req openAIInputTokensCountRequest
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
return nil, fmt.Errorf("parse responses input_tokens request: %w", err)
|
||||
}
|
||||
originalModel := strings.TrimSpace(req.Model)
|
||||
if originalModel == "" {
|
||||
return nil, fmt.Errorf("parse responses input_tokens request: model is required")
|
||||
}
|
||||
billingModel := resolveOpenAIForwardModel(account, originalModel, "")
|
||||
upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel)
|
||||
req.Model = upstreamModel
|
||||
return &openAIInputTokensCountPrepared{
|
||||
Request: req,
|
||||
OriginalModel: originalModel,
|
||||
NormalizedModel: originalModel,
|
||||
BillingModel: billingModel,
|
||||
UpstreamModel: upstreamModel,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func shouldEstimateOpenAIInputTokensLocally(account *Account) bool {
|
||||
if account == nil || account.IsGrok() || account.IsCNProvider() || account.Type == AccountTypeUpstream {
|
||||
return true
|
||||
}
|
||||
if account.Type != AccountTypeAPIKey {
|
||||
return false
|
||||
}
|
||||
rawBaseURL := strings.TrimSpace(account.GetCredential("base_url"))
|
||||
if rawBaseURL == "" {
|
||||
return false
|
||||
}
|
||||
parsed, err := url.Parse(rawBaseURL)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return !strings.EqualFold(parsed.Hostname(), "api.openai.com")
|
||||
}
|
||||
|
||||
func isOpenAIResponsesInputTokensUnsupported(account *Account, statusCode int, body []byte) bool {
|
||||
if statusCode == http.StatusNotFound {
|
||||
return true
|
||||
}
|
||||
return account != nil && account.Type == AccountTypeOAuth && isOpenAIOAuthInputTokensUnsupported(statusCode, body)
|
||||
}
|
||||
|
||||
func writeOpenAIResponsesInputTokensFallback(c *gin.Context, account *Account, prepared *openAIInputTokensCountPrepared, statusCode int, reason string) {
|
||||
estimated := openAIInputTokensFallbackMinimum
|
||||
if prepared != nil {
|
||||
if got, err := estimateOpenAIInputTokens(prepared.Request); err == nil && got > 0 {
|
||||
estimated = got
|
||||
}
|
||||
}
|
||||
accountID := int64(0)
|
||||
upstreamModel := ""
|
||||
if account != nil {
|
||||
accountID = account.ID
|
||||
}
|
||||
if prepared != nil {
|
||||
upstreamModel = prepared.UpstreamModel
|
||||
}
|
||||
logger.L().Info("openai responses input_tokens: local estimate fallback",
|
||||
zap.Int64("account_id", accountID),
|
||||
zap.Int("upstream_status", statusCode),
|
||||
zap.Int("estimated_input_tokens", estimated),
|
||||
zap.String("upstream_model", upstreamModel),
|
||||
zap.String("reason", reason),
|
||||
)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"object": "response.input_tokens",
|
||||
"input_tokens": estimated,
|
||||
})
|
||||
}
|
||||
|
||||
func writeOpenAIResponsesInputTokensError(c *gin.Context, status int, errType, message string) {
|
||||
c.JSON(status, gin.H{
|
||||
"error": gin.H{
|
||||
"type": errType,
|
||||
"message": message,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// EstimateGrokCountTokens estimates an Anthropic-compatible count_tokens request
|
||||
// locally. Grok does not expose a compatible token-counting endpoint, so this
|
||||
// path deliberately avoids account selection, credentials, and upstream calls.
|
||||
|
||||
@@ -420,6 +420,12 @@ func IsForwardableOpenAIResponsesRequestPath(c *gin.Context) bool {
|
||||
return ok
|
||||
}
|
||||
|
||||
// IsOpenAIResponsesInputTokensRequestPath reports whether the request targets
|
||||
// the native Responses input-token counting endpoint.
|
||||
func IsOpenAIResponsesInputTokensRequestPath(c *gin.Context) bool {
|
||||
return openAIResponsesRequestPathSuffix(c) == "/input_tokens"
|
||||
}
|
||||
|
||||
// rawOpenAIResponsesRequestPathSuffix 仅做提取,不做任何安全判断。
|
||||
func rawOpenAIResponsesRequestPathSuffix(c *gin.Context) string {
|
||||
if c == nil || c.Request == nil || c.Request.URL == nil {
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
func TestForwardResponsesInputTokensCustomRelayUsesLocalEstimate(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/input_tokens", nil)
|
||||
|
||||
upstream := &httpUpstreamRecorder{}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
account := &Account{
|
||||
ID: 159,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "relay-key",
|
||||
"base_url": "https://relay.example/v1",
|
||||
},
|
||||
}
|
||||
body := []byte(`{"model":"gpt-5.4","instructions":"Be concise.","input":"hello world","tools":[{"type":"function","name":"lookup","description":"Look up a value","parameters":{"type":"object"}}]}`)
|
||||
|
||||
err := svc.ForwardResponsesInputTokens(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.Equal(t, "response.input_tokens", gjson.Get(recorder.Body.String(), "object").String())
|
||||
require.Positive(t, gjson.Get(recorder.Body.String(), "input_tokens").Int())
|
||||
require.Nil(t, upstream.lastReq, "custom relay must not receive /v1/responses/input_tokens")
|
||||
}
|
||||
|
||||
func TestForwardResponsesInputTokensGrokOAuthUsesLocalEstimate(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/input_tokens", nil)
|
||||
|
||||
upstream := &httpUpstreamRecorder{}
|
||||
svc := &OpenAIGatewayService{httpUpstream: upstream}
|
||||
account := &Account{ID: 160, Platform: PlatformGrok, Type: AccountTypeOAuth}
|
||||
body := []byte(`{"model":"grok-4.1","input":"hello world"}`)
|
||||
|
||||
err := svc.ForwardResponsesInputTokens(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.Equal(t, "response.input_tokens", gjson.Get(recorder.Body.String(), "object").String())
|
||||
require.Positive(t, gjson.Get(recorder.Body.String(), "input_tokens").Int())
|
||||
require.Nil(t, upstream.lastReq)
|
||||
}
|
||||
|
||||
func TestForwardResponsesInputTokensUpstream404FallsBackLocally(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/input_tokens", nil)
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusNotFound,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(`{"error":{"type":"invalid_request_error","message":"Invalid URL (POST /v1/responses/input_tokens)"}}`)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
account := &Account{
|
||||
ID: 171,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "official-key",
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
},
|
||||
}
|
||||
body := []byte(`{"model":"gpt-5.4","instructions":"Be concise.","input":"hello world"}`)
|
||||
|
||||
err := svc.ForwardResponsesInputTokens(context.Background(), c, account, body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.Equal(t, "response.input_tokens", gjson.Get(recorder.Body.String(), "object").String())
|
||||
require.Positive(t, gjson.Get(recorder.Body.String(), "input_tokens").Int())
|
||||
require.NotNil(t, upstream.lastReq)
|
||||
}
|
||||
@@ -134,6 +134,7 @@ func TestOpenAIResponsesRequestPathSuffixRejectsNonConformingSubpaths(t *testing
|
||||
for path, want := range map[string]string{
|
||||
"/v1/responses": "",
|
||||
"/v1/responses/compact": "/compact",
|
||||
"/v1/responses/input_tokens": "/input_tokens",
|
||||
"/responses/compact/": "/compact",
|
||||
"/backend-api/codex/responses/compact": "/compact",
|
||||
} {
|
||||
@@ -145,6 +146,15 @@ func TestOpenAIResponsesRequestPathSuffixRejectsNonConformingSubpaths(t *testing
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsOpenAIResponsesInputTokensRequestPath(t *testing.T) {
|
||||
for _, path := range []string{"/v1/responses/input_tokens", "/responses/input_tokens", "/backend-api/codex/responses/input_tokens"} {
|
||||
c := newResponsesSuffixTestContext(t, path)
|
||||
require.True(t, IsOpenAIResponsesInputTokensRequestPath(c), "path=%s", path)
|
||||
}
|
||||
c := newResponsesSuffixTestContext(t, "/v1/responses/compact")
|
||||
require.False(t, IsOpenAIResponsesInputTokensRequestPath(c))
|
||||
}
|
||||
|
||||
func TestIsOpenAIResponsesCompactPathUsesLegacyEndpointShape(t *testing.T) {
|
||||
legacyPaths := []string{
|
||||
"/v1/responses/compact",
|
||||
|
||||
Reference in New Issue
Block a user