Bridge OpenAI count_tokens to responses input_tokens

This commit is contained in:
JRBaggins
2026-06-26 17:59:26 +08:00
parent df99b9417c
commit 7a38c66214
7 changed files with 512 additions and 3 deletions
@@ -0,0 +1,145 @@
package handler
import (
"net/http"
"strconv"
"time"
"github.com/Wei-Shaw/sub2api/internal/domain"
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
// CountTokens handles Anthropic-compatible POST /v1/messages/count_tokens for OpenAI groups.
// It validates billing and routes to an OpenAI token-count bridge without taking concurrency slots
// or recording usage.
func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok {
h.anthropicErrorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
return
}
subject, ok := middleware2.GetAuthSubjectFromContext(c)
if !ok {
h.anthropicErrorResponse(c, http.StatusInternalServerError, "api_error", "User context not found")
return
}
reqLog := requestLogger(
c,
"handler.openai_gateway.count_tokens",
zap.Int64("user_id", subject.UserID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
)
if apiKey.Group != nil && !apiKey.Group.AllowMessagesDispatch {
h.anthropicErrorResponse(c, http.StatusForbidden, "permission_error",
"This group does not allow /v1/messages dispatch")
return
}
if !h.ensureResponsesDependencies(c, reqLog) {
return
}
body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
if err != nil {
if maxErr, ok := extractMaxBytesError(err); ok {
h.anthropicErrorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
return
}
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
return
}
if len(body) == 0 {
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
return
}
bodyRef := service.NewRequestBodyRef(body)
parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic)
if err != nil {
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
}
if parsedReq.Model == "" {
h.anthropicErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
return
}
reqModel := parsedReq.Model
routingModel := service.NormalizeOpenAICompatRequestedModel(reqModel)
preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(apiKey, reqModel)
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", parsedReq.Stream))
setOpsRequestContext(c, reqModel, false)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(false, false)))
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel)
mappedBodyForMessages := newOpenAIModelMappedBodyCache(body, h.gatewayService.ReplaceModelInBody)
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_count_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.anthropicErrorResponse(c, status, code, message)
return
}
requestStart := time.Now()
sessionHash := h.gatewayService.GenerateSessionHash(c, body)
currentRoutingModel := routingModel
if preferredMappedModel != "" {
currentRoutingModel = preferredMappedModel
}
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
"",
sessionHash,
currentRoutingModel,
nil,
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
openAICompatibleRequestPlatform(apiKey),
)
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)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
}
h.anthropicErrorResponse(c, cls.Status, cls.ErrType, cls.Message)
return
}
if selection == nil || selection.Account == nil {
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel, service.PlatformOpenAI)
if !cls.ModelNotFound {
markOpsRoutingCapacityLimited(c)
}
h.anthropicErrorResponse(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()
}
forwardBody := mappedBodyForMessages(channelMapping.Mapped, channelMapping.MappedModel)
defaultMappedModel := preferredMappedModel
if err := h.gatewayService.ForwardCountTokensAsAnthropic(c.Request.Context(), c, account, forwardBody, defaultMappedModel); err != nil {
reqLog.Error("openai_count_tokens.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err))
}
}
+6 -1
View File
@@ -73,8 +73,13 @@ func RegisterGatewayRoutes(
}
h.Gateway.Messages(c)
})
// /v1/messages/count_tokens: OpenAI groups get 404
// /v1/messages/count_tokens: OpenAI uses Anthropic-compat bridge; other
// OpenAI-compatible platforms keep the prior unsupported response.
gateway.POST("/messages/count_tokens", func(c *gin.Context) {
if isOpenAIGatewayPlatform(c) {
h.OpenAIGateway.CountTokens(c)
return
}
if isOpenAIResponsesCompatibleGatewayPlatform(c) {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusNotFound, gin.H{
@@ -106,6 +106,14 @@ func TestGatewayRoutesGrokOnlyAllowsResponsesHTTP(t *testing.T) {
require.Contains(t, w.Body.String(), "not supported for Grok groups")
}
req := httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", strings.NewReader(`{"model":"grok","messages":[{"role":"user","content":"hi"}]}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.Equal(t, http.StatusNotFound, w.Code)
require.Contains(t, w.Body.String(), "Token counting is not supported for this platform")
for _, path := range []string{
"/v1/responses",
"/responses",
@@ -119,3 +127,14 @@ func TestGatewayRoutesGrokOnlyAllowsResponsesHTTP(t *testing.T) {
require.NotEqual(t, http.StatusNotFound, w.Code, "path=%s should still reach Responses handler", path)
}
}
func TestGatewayRoutesOpenAICountTokensPathIsRegistered(t *testing.T) {
router := newGatewayRoutesTestRouter(service.PlatformOpenAI)
req := httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", strings.NewReader(`{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"hi"}]}`))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
require.NotEqual(t, http.StatusNotFound, w.Code)
}
@@ -18,6 +18,10 @@ func buildOpenAIEndpointURL(base string, endpoint string) string {
return normalized + endpoint
}
func buildOpenAIResponsesInputTokensURL(base string) string {
return buildOpenAIEndpointURL(base, "/v1/responses/input_tokens")
}
func openAIBaseURLHasVersionSuffix(raw string) bool {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
@@ -0,0 +1,228 @@
package service
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"go.uber.org/zap"
)
type openAIInputTokensCountRequest struct {
Model string `json:"model"`
Instructions string `json:"instructions,omitempty"`
Input json.RawMessage `json:"input,omitempty"`
Tools []apicompat.ResponsesTool `json:"tools,omitempty"`
ToolChoice json.RawMessage `json:"tool_choice,omitempty"`
}
// ForwardCountTokensAsAnthropic bridges Anthropic /v1/messages/count_tokens to
// OpenAI POST /v1/responses/input_tokens and returns Anthropic-compatible output.
func (s *OpenAIGatewayService) ForwardCountTokensAsAnthropic(
ctx context.Context,
c *gin.Context,
account *Account,
body []byte,
defaultMappedModel string,
) error {
if account == nil {
writeAnthropicCountTokensError(c, http.StatusServiceUnavailable, "api_error", "No available OpenAI accounts")
return fmt.Errorf("count_tokens: missing account")
}
var anthropicReq apicompat.AnthropicRequest
if err := json.Unmarshal(body, &anthropicReq); err != nil {
writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return fmt.Errorf("parse anthropic count_tokens request: %w", err)
}
originalModel := anthropicReq.Model
applyOpenAICompatModelNormalization(&anthropicReq)
normalizedModel := anthropicReq.Model
billingModel := resolveOpenAIForwardModel(account, normalizedModel, strings.TrimSpace(defaultMappedModel))
upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel)
responsesReq, err := apicompat.AnthropicToResponses(&anthropicReq)
if err != nil {
writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "Failed to convert request body")
return fmt.Errorf("convert anthropic request to responses: %w", err)
}
upstreamBody, err := marshalOpenAIUpstreamJSON(openAIInputTokensCountRequest{
Model: upstreamModel,
Instructions: responsesReq.Instructions,
Input: responsesReq.Input,
Tools: responsesReq.Tools,
ToolChoice: responsesReq.ToolChoice,
})
if err != nil {
writeAnthropicCountTokensError(c, http.StatusInternalServerError, "api_error", "Failed to build request")
return fmt.Errorf("marshal openai input_tokens body: %w", err)
}
logger.L().Debug("openai count_tokens: model mapping applied",
zap.Int64("account_id", account.ID),
zap.String("original_model", originalModel),
zap.String("normalized_model", normalizedModel),
zap.String("billing_model", billingModel),
zap.String("upstream_model", upstreamModel),
)
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Failed to get access token")
return fmt.Errorf("get access token: %w", err)
}
upstreamReq, err := s.buildInputTokensUpstreamRequest(ctx, c, account, upstreamBody, token)
if err != nil {
writeAnthropicCountTokensError(c, http.StatusInternalServerError, "api_error", "Failed to build request")
return fmt.Errorf("build input_tokens 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, "")
writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
return fmt.Errorf("openai input_tokens upstream request failed: %s", safeErr)
}
defer func() { _ = resp.Body.Close() }()
respBody, err := io.ReadAll(resp.Body)
if err != nil {
writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Failed to read response")
return fmt.Errorf("read input_tokens response: %w", err)
}
if resp.StatusCode >= 400 {
if s.rateLimitService != nil {
s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
}
upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
if account.Type == AccountTypeOAuth && isOpenAIOAuthInputTokensUnsupported(resp.StatusCode) {
writeAnthropicCountTokensError(c, http.StatusNotFound, "not_found_error", "Token counting is not supported for this OpenAI account type")
return nil
}
if isOpenAIInputTokensUnsupported(resp.StatusCode, respBody) {
writeAnthropicCountTokensError(c, http.StatusNotFound, "not_found_error", "Token counting is not supported by upstream")
return nil
}
upstreamDetail := ""
if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody {
maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes
if maxBytes <= 0 {
maxBytes = 2048
}
upstreamDetail = truncateString(string(respBody), maxBytes)
}
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail)
errMsg := "Upstream request failed"
switch resp.StatusCode {
case 429:
errMsg = "Rate limit exceeded"
case 500, 502, 503, 504, 529:
errMsg = "Upstream service temporarily unavailable"
}
writeAnthropicCountTokensError(c, resp.StatusCode, "upstream_error", errMsg)
if upstreamMsg == "" {
return fmt.Errorf("input_tokens upstream error: %d", resp.StatusCode)
}
return fmt.Errorf("input_tokens upstream error: %d message=%s", resp.StatusCode, upstreamMsg)
}
inputTokens := gjson.GetBytes(respBody, "input_tokens")
if !inputTokens.Exists() {
writeAnthropicCountTokensError(c, http.StatusBadGateway, "upstream_error", "Upstream response missing input_tokens")
return fmt.Errorf("input_tokens response missing input_tokens field")
}
c.JSON(http.StatusOK, gin.H{
"input_tokens": int(inputTokens.Int()),
})
return nil
}
func (s *OpenAIGatewayService) buildInputTokensUpstreamRequest(
ctx context.Context,
c *gin.Context,
account *Account,
body []byte,
token string,
) (*http.Request, error) {
targetURL := openaiPlatformAPIInputTokensURL
if account.Type == AccountTypeAPIKey {
if baseURL := account.GetOpenAIBaseURL(); strings.TrimSpace(baseURL) != "" {
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
return nil, err
}
targetURL = buildOpenAIResponsesInputTokensURL(validatedURL)
}
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body))
if err != nil {
return nil, err
}
req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI))
req.Header.Set("authorization", "Bearer "+token)
req.Header.Set("content-type", "application/json")
req.Header.Set("accept", "application/json")
if c != nil && c.Request != nil {
for key, values := range c.Request.Header {
lower := strings.ToLower(strings.TrimSpace(key))
if lower != "user-agent" && lower != "accept-language" {
continue
}
for _, v := range values {
req.Header.Add(key, v)
}
}
}
return req, nil
}
func writeAnthropicCountTokensError(c *gin.Context, status int, errType, message string) {
c.JSON(status, gin.H{
"type": "error",
"error": gin.H{
"type": errType,
"message": message,
},
})
}
func isOpenAIInputTokensUnsupported(statusCode int, body []byte) bool {
if statusCode != http.StatusNotFound {
return false
}
msg := strings.ToLower(strings.TrimSpace(extractUpstreamErrorMessage(body)))
return strings.Contains(msg, "input_tokens") && strings.Contains(msg, "not found")
}
func isOpenAIOAuthInputTokensUnsupported(statusCode int) bool {
switch statusCode {
case http.StatusUnauthorized, http.StatusForbidden, http.StatusNotFound:
return true
default:
return false
}
}
@@ -0,0 +1,107 @@
package service
import (
"bytes"
"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 TestOpenAIGatewayService_ForwardCountTokensAsAnthropic_APIKeyUsesResponsesInputTokens(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"claude-sonnet-4-5","system":"You are helpful.","messages":[{"role":"user","content":"hello"}],"tools":[{"name":"lookup","input_schema":{"type":"object"}}]}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"object":"response.input_tokens","input_tokens":42}`)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{
Enabled: false,
AllowInsecureHTTP: true,
}}},
httpUpstream: upstream,
}
account := &Account{
ID: 101,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "http://upstream.example",
},
Status: StatusActive,
Schedulable: true,
}
err := svc.ForwardCountTokensAsAnthropic(context.Background(), c, account, body, "gpt-5.3-codex")
require.NoError(t, err)
require.Equal(t, http.StatusOK, rec.Code)
require.JSONEq(t, `{"input_tokens":42}`, rec.Body.String())
require.NotNil(t, upstream.lastReq)
require.Equal(t, "http://upstream.example/v1/responses/input_tokens", upstream.lastReq.URL.String())
require.Equal(t, "Bearer sk-test", upstream.lastReq.Header.Get("authorization"))
require.Equal(t, "gpt-5.3-codex", gjson.GetBytes(upstream.lastBody, "model").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "input").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "messages").Exists())
}
func TestOpenAIGatewayService_ForwardCountTokensAsAnthropic_OAuthFallsBackWhenPlatformEndpointUnsupported(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"claude-opus-4-1","messages":[{"role":"user","content":"hello"}]}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Header.Set("User-Agent", "Claude-Code/1.0")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusUnauthorized,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"type":"invalid_request_error","message":"unauthorized"}}`)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{},
httpUpstream: upstream,
}
account := &Account{
ID: 202,
Name: "openai-oauth",
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "oauth-token",
},
Status: StatusActive,
Schedulable: true,
}
err := svc.ForwardCountTokensAsAnthropic(context.Background(), c, account, body, "gpt-5.4")
require.NoError(t, err)
require.Equal(t, http.StatusNotFound, rec.Code)
require.Contains(t, rec.Body.String(), "Token counting is not supported for this OpenAI account type")
require.NotNil(t, upstream.lastReq)
require.Equal(t, "https://api.openai.com/v1/responses/input_tokens", upstream.lastReq.URL.String())
require.Equal(t, "Bearer oauth-token", upstream.lastReq.Header.Get("authorization"))
require.Empty(t, upstream.lastReq.Header.Get("Chatgpt-Account-Id"))
}
@@ -41,8 +41,9 @@ const (
// ChatGPT internal API for OAuth accounts
chatgptCodexURL = "https://chatgpt.com/backend-api/codex/responses"
// OpenAI Platform API for API Key accounts (fallback)
openaiPlatformAPIURL = "https://api.openai.com/v1/responses"
openaiStickySessionTTL = time.Hour // 粘性会话TTL
openaiPlatformAPIURL = "https://api.openai.com/v1/responses"
openaiPlatformAPIInputTokensURL = "https://api.openai.com/v1/responses/input_tokens"
openaiStickySessionTTL = time.Hour // 粘性会话TTL
// 与真实 Codex CLI 的 User-Agent 结构对齐:
// {originator}/{version} ({OS} {OS_version}; {arch}) {terminal}
// 旧值 "codex_cli_rs/0.125.0" 缺少 OS/架构/终端后缀,易被上游指纹识别为非官方客户端。