diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index df8776981c..6985392926 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -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. diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go index 3f7b994702..005c397d44 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -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 diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index 70e61e24f1..f300b31b95 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -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. diff --git a/backend/internal/handler/openai_responses_input_tokens_test.go b/backend/internal/handler/openai_responses_input_tokens_test.go new file mode 100644 index 0000000000..b3260fc2ad --- /dev/null +++ b/backend/internal/handler/openai_responses_input_tokens_test.go @@ -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")) +} diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index 8be450b16f..a5343d9077 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -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 } diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 9a06ae9c79..de24061db3 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -170,6 +170,10 @@ func RegisterGatewayRoutes( }) return } + if service.IsOpenAIResponsesInputTokensRequestPath(c) && isOpenAIResponsesCompatibleGatewayPlatform(c) { + h.OpenAIGateway.ResponsesInputTokens(c) + return + } next(c) } } diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index bcdaa457b8..2cd4519d07 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -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. diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 3559fb400d..f2efddcf77 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -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 { diff --git a/backend/internal/service/openai_responses_input_tokens_test.go b/backend/internal/service/openai_responses_input_tokens_test.go new file mode 100644 index 0000000000..5258d4bb70 --- /dev/null +++ b/backend/internal/service/openai_responses_input_tokens_test.go @@ -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) +} diff --git a/backend/internal/service/upstream_path_guard_test.go b/backend/internal/service/upstream_path_guard_test.go index e04781e6a7..f915ecd394 100644 --- a/backend/internal/service/upstream_path_guard_test.go +++ b/backend/internal/service/upstream_path_guard_test.go @@ -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",