Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
This commit is contained in:
@@ -0,0 +1,777 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"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"
|
||||
"github.com/tiktoken-go/tokenizer"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const (
|
||||
openAIResponsesInputItemTokenOverhead = 3
|
||||
openAIResponsesContentPartOverhead = 1
|
||||
openAIInputTokensFallbackMinimum = 1
|
||||
)
|
||||
|
||||
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"`
|
||||
}
|
||||
|
||||
type openAIInputTokensCountPrepared struct {
|
||||
Request openAIInputTokensCountRequest
|
||||
OriginalModel string
|
||||
NormalizedModel string
|
||||
BillingModel string
|
||||
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.
|
||||
func EstimateGrokCountTokens(body []byte) (int, error) {
|
||||
return estimateAnthropicCountTokensLocally(body)
|
||||
}
|
||||
|
||||
// estimateAnthropicCountTokensLocally 走 Anthropic→Responses→tiktoken 链本地估算
|
||||
// count_tokens,不发任何上游请求(上游无兼容端点的平台使用)。
|
||||
func estimateAnthropicCountTokensLocally(body []byte) (int, error) {
|
||||
var anthropicReq apicompat.AnthropicRequest
|
||||
if err := json.Unmarshal(body, &anthropicReq); err != nil {
|
||||
return 0, fmt.Errorf("parse anthropic count_tokens request: %w", err)
|
||||
}
|
||||
if strings.TrimSpace(anthropicReq.Model) == "" {
|
||||
return 0, fmt.Errorf("parse anthropic count_tokens request: model is required")
|
||||
}
|
||||
|
||||
responsesReq, err := apicompat.AnthropicToResponses(&anthropicReq)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("convert anthropic request to responses: %w", err)
|
||||
}
|
||||
|
||||
estimated, err := estimateOpenAIInputTokens(openAIInputTokensCountRequest{
|
||||
Model: anthropicReq.Model,
|
||||
Instructions: responsesReq.Instructions,
|
||||
Input: responsesReq.Input,
|
||||
Tools: responsesReq.Tools,
|
||||
ToolChoice: responsesReq.ToolChoice,
|
||||
})
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("estimate input tokens: %w", err)
|
||||
}
|
||||
if estimated < openAIInputTokensFallbackMinimum {
|
||||
estimated = openAIInputTokensFallbackMinimum
|
||||
}
|
||||
return estimated, nil
|
||||
}
|
||||
|
||||
// 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")
|
||||
}
|
||||
|
||||
// 国产供应商(全部协议,含 anthropic):一律本地估算,不发上游请求。
|
||||
// 依据(2026-08 核实):三家的 Anthropic 兼容层均未提供
|
||||
// /v1/messages/count_tokens——DeepSeek 官方 anthropic_api 文档无此端点
|
||||
// (且注明 anthropic-version 头被忽略),聚合网关 OpenModel 明确标注
|
||||
// count_tokens 为 "Anthropic only",Kimi/智谱亦无任何文档承诺。转发上游
|
||||
// 只会常态 404,且错误还会流入账号处置逻辑误伤整账号调度;Claude Code
|
||||
// 高频调用此端点,本地 tiktoken 估算是与 Grok 一致的既有方案。
|
||||
if account.IsCNProvider() {
|
||||
estimated, err := estimateAnthropicCountTokensLocally(body)
|
||||
if err != nil {
|
||||
writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
|
||||
return fmt.Errorf("count_tokens: estimate cn provider input tokens: %w", err)
|
||||
}
|
||||
logger.L().Debug("openai count_tokens: cn provider local estimate",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Int("estimated_input_tokens", estimated),
|
||||
)
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"input_tokens": estimated,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
prepared, err := prepareOpenAIInputTokensCountRequest(body, account, defaultMappedModel)
|
||||
if err != nil {
|
||||
writeAnthropicCountTokensError(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
|
||||
return err
|
||||
}
|
||||
|
||||
upstreamBody, err := marshalOpenAIUpstreamJSON(prepared.Request)
|
||||
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", prepared.OriginalModel),
|
||||
zap.String("normalized_model", prepared.NormalizedModel),
|
||||
zap.String("billing_model", prepared.BillingModel),
|
||||
zap.String("upstream_model", prepared.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 {
|
||||
upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
|
||||
if account.Type == AccountTypeOAuth && isOpenAIOAuthInputTokensUnsupported(resp.StatusCode, respBody) {
|
||||
writeOpenAIOAuthInputTokensFallback(c, account, prepared, resp.StatusCode)
|
||||
return nil
|
||||
}
|
||||
|
||||
if s.rateLimitService != nil {
|
||||
s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
|
||||
}
|
||||
|
||||
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 prepareOpenAIInputTokensCountRequest(
|
||||
body []byte,
|
||||
account *Account,
|
||||
defaultMappedModel string,
|
||||
) (*openAIInputTokensCountPrepared, error) {
|
||||
var anthropicReq apicompat.AnthropicRequest
|
||||
if err := json.Unmarshal(body, &anthropicReq); err != nil {
|
||||
return nil, 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 {
|
||||
return nil, fmt.Errorf("convert anthropic request to responses: %w", err)
|
||||
}
|
||||
|
||||
return &openAIInputTokensCountPrepared{
|
||||
Request: openAIInputTokensCountRequest{
|
||||
Model: upstreamModel,
|
||||
Instructions: responsesReq.Instructions,
|
||||
Input: responsesReq.Input,
|
||||
Tools: responsesReq.Tools,
|
||||
ToolChoice: responsesReq.ToolChoice,
|
||||
},
|
||||
OriginalModel: originalModel,
|
||||
NormalizedModel: normalizedModel,
|
||||
BillingModel: billingModel,
|
||||
UpstreamModel: upstreamModel,
|
||||
}, 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))
|
||||
authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for key, values := range authHeaders {
|
||||
for _, value := range values {
|
||||
req.Header.Add(key, value)
|
||||
}
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)
|
||||
account.ApplyHeaderOverrides(req.Header)
|
||||
|
||||
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 writeOpenAIOAuthInputTokensFallback(c *gin.Context, account *Account, prepared *openAIInputTokensCountPrepared, statusCode int) {
|
||||
estimated := openAIInputTokensFallbackMinimum
|
||||
if got, err := estimateOpenAIInputTokens(prepared.Request); err == nil {
|
||||
if got > 0 {
|
||||
estimated = got
|
||||
}
|
||||
logger.L().Info("openai count_tokens: oauth fallback to local tiktoken estimate",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Int("upstream_status", statusCode),
|
||||
zap.Int("estimated_input_tokens", estimated),
|
||||
zap.String("upstream_model", prepared.UpstreamModel),
|
||||
)
|
||||
} else {
|
||||
logger.L().Warn("openai count_tokens: oauth local tiktoken fallback failed, using minimum estimate",
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.Int("upstream_status", statusCode),
|
||||
zap.Int("estimated_input_tokens", estimated),
|
||||
zap.String("upstream_model", prepared.UpstreamModel),
|
||||
zap.Error(err),
|
||||
)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"input_tokens": estimated,
|
||||
})
|
||||
}
|
||||
|
||||
func isOpenAIOAuthInputTokensUnsupported(statusCode int, body []byte) bool {
|
||||
switch statusCode {
|
||||
case http.StatusUnauthorized, http.StatusForbidden, http.StatusNotFound:
|
||||
default:
|
||||
return false
|
||||
}
|
||||
|
||||
bodyLower := strings.ToLower(string(body))
|
||||
msg := strings.ToLower(strings.TrimSpace(extractUpstreamErrorMessage(body)))
|
||||
code := strings.ToLower(strings.TrimSpace(extractUpstreamErrorCode(body)))
|
||||
|
||||
if code == "missing_scope" ||
|
||||
strings.Contains(bodyLower, "api.responses.write") ||
|
||||
strings.Contains(bodyLower, "missing scopes") ||
|
||||
strings.Contains(bodyLower, "insufficient_scope") {
|
||||
return true
|
||||
}
|
||||
|
||||
if statusCode == http.StatusNotFound && isOpenAIInputTokensUnsupported(statusCode, body) {
|
||||
return true
|
||||
}
|
||||
|
||||
// OAuth's platform endpoint can be blocked by an upstream proxy before it
|
||||
// reaches the API and return an HTML 403 page without a structured error.
|
||||
// Treat that endpoint-level response like the other unsupported cases so
|
||||
// count_tokens remains a local, non-health-affecting convenience request.
|
||||
if statusCode == http.StatusForbidden && isHTMLResponse(body) {
|
||||
return true
|
||||
}
|
||||
|
||||
return strings.Contains(msg, "input_tokens") &&
|
||||
(strings.Contains(msg, "not found") ||
|
||||
strings.Contains(msg, "not supported") ||
|
||||
strings.Contains(msg, "unsupported"))
|
||||
}
|
||||
|
||||
func isHTMLResponse(body []byte) bool {
|
||||
trimmed := strings.TrimSpace(strings.ToLower(string(body)))
|
||||
return strings.HasPrefix(trimmed, "<!doctype html") ||
|
||||
strings.HasPrefix(trimmed, "<html")
|
||||
}
|
||||
|
||||
func estimateOpenAIInputTokens(req openAIInputTokensCountRequest) (int, error) {
|
||||
codec, err := openAIInputTokensCodecForModel(req.Model)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
total := 0
|
||||
addCount := func(text string) error {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
n, err := codec.Count(text)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
total += n
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := addCount(req.Instructions); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
inputTokens, err := estimateOpenAIInputTokensForInput(codec, req.Input)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
total += inputTokens
|
||||
|
||||
for _, tool := range req.Tools {
|
||||
raw, err := marshalOpenAIUpstreamJSON(tool)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := addCount(string(raw)); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
if len(req.ToolChoice) > 0 {
|
||||
compacted, err := compactOpenAIInputTokensJSON(req.ToolChoice)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := addCount(compacted); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
|
||||
if total < 0 {
|
||||
return 0, nil
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func estimateOpenAIInputTokensForInput(codec tokenizer.Codec, raw json.RawMessage) (int, error) {
|
||||
if len(bytes.TrimSpace(raw)) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
var plainText string
|
||||
if err := json.Unmarshal(raw, &plainText); err == nil {
|
||||
return codec.Count(plainText)
|
||||
}
|
||||
|
||||
var items []apicompat.ResponsesInputItem
|
||||
if err := json.Unmarshal(raw, &items); err == nil {
|
||||
return estimateOpenAIInputTokensForInputItems(codec, items)
|
||||
}
|
||||
|
||||
compacted, err := compactOpenAIInputTokensJSON(raw)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return codec.Count(compacted)
|
||||
}
|
||||
|
||||
func estimateOpenAIInputTokensForInputItems(codec tokenizer.Codec, items []apicompat.ResponsesInputItem) (int, error) {
|
||||
total := 0
|
||||
countText := func(text string) error {
|
||||
text = strings.TrimSpace(text)
|
||||
if text == "" {
|
||||
return nil
|
||||
}
|
||||
n, err := codec.Count(text)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
total += n
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, item := range items {
|
||||
total += openAIResponsesInputItemTokenOverhead
|
||||
if err := countText(item.Role); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if item.Type != "" && item.Type != "message" {
|
||||
if err := countText(item.Type); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
if err := countText(item.Name); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := countText(item.Arguments); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := countText(item.Output); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := countText(item.CallID); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := countText(item.ID); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
if len(bytes.TrimSpace(item.Content)) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
var contentText string
|
||||
if err := json.Unmarshal(item.Content, &contentText); err == nil {
|
||||
if err := countText(contentText); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
var parts []apicompat.ResponsesContentPart
|
||||
if err := json.Unmarshal(item.Content, &parts); err == nil {
|
||||
for _, part := range parts {
|
||||
total += openAIResponsesContentPartOverhead
|
||||
switch part.Type {
|
||||
case "input_text", "output_text", "text":
|
||||
if err := countText(part.Text); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
case "input_image":
|
||||
if err := countText(estimateOpenAIInputImageText(part.ImageURL)); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
default:
|
||||
if err := countText(part.Type); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
compacted, err := compactOpenAIInputTokensJSON(item.Content)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := countText(compacted); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func estimateOpenAIInputImageText(imageURL string) string {
|
||||
trimmed := strings.TrimSpace(imageURL)
|
||||
if trimmed == "" {
|
||||
return ""
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(trimmed), "data:") {
|
||||
if comma := strings.Index(trimmed, ","); comma > 0 {
|
||||
return trimmed[:comma]
|
||||
}
|
||||
}
|
||||
return trimmed
|
||||
}
|
||||
|
||||
func compactOpenAIInputTokensJSON(raw json.RawMessage) (string, error) {
|
||||
if len(bytes.TrimSpace(raw)) == 0 {
|
||||
return "", nil
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := json.Compact(&buf, raw); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
func openAIInputTokensCodecForModel(model string) (tokenizer.Codec, error) {
|
||||
switch openAIInputTokensEncodingForModel(model) {
|
||||
case tokenizer.Cl100kBase:
|
||||
return tokenizer.Get(tokenizer.Cl100kBase)
|
||||
default:
|
||||
return tokenizer.Get(tokenizer.O200kBase)
|
||||
}
|
||||
}
|
||||
|
||||
func openAIInputTokensEncodingForModel(model string) tokenizer.Encoding {
|
||||
normalized := strings.ToLower(strings.TrimSpace(model))
|
||||
switch {
|
||||
case strings.HasPrefix(normalized, "gpt-3.5"),
|
||||
(strings.HasPrefix(normalized, "gpt-4") &&
|
||||
!strings.HasPrefix(normalized, "gpt-4o") &&
|
||||
!strings.HasPrefix(normalized, "gpt-4.1")),
|
||||
strings.HasPrefix(normalized, "text-embedding-"):
|
||||
return tokenizer.Cl100kBase
|
||||
default:
|
||||
return tokenizer.O200kBase
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user