mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
Merge remote-tracking branch 'origin/main' into fix/issue-5796-composite-new-platforms
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
|
||||
|
||||
@@ -30,8 +30,8 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
|
||||
return
|
||||
}
|
||||
if apiKey.Group.Platform != service.PlatformOpenAI {
|
||||
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex alpha search is only available for OpenAI groups")
|
||||
if apiKey.Group.Platform != service.PlatformOpenAI && apiKey.Group.Platform != service.PlatformComposite {
|
||||
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex alpha search is only available for OpenAI and Composite groups")
|
||||
return
|
||||
}
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
@@ -75,6 +75,10 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
requestedModel := strings.TrimSpace(modelResult.String())
|
||||
if !compositeTargetPlatformAllowed(c, apiKey, requestedModel, service.PlatformOpenAI) {
|
||||
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex alpha search only supports OpenAI models for Composite groups")
|
||||
return
|
||||
}
|
||||
reqLog = reqLog.With(zap.String("model", requestedModel))
|
||||
setOpsRequestContext(c, requestedModel, false)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
|
||||
|
||||
@@ -27,8 +27,8 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
|
||||
h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required")
|
||||
return
|
||||
}
|
||||
if apiKey.Group.Platform != service.PlatformOpenAI {
|
||||
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex models manifest is only available for OpenAI groups")
|
||||
if apiKey.Group.Platform != service.PlatformOpenAI && apiKey.Group.Platform != service.PlatformComposite {
|
||||
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex models manifest is only available for OpenAI and Composite groups")
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -116,6 +116,19 @@ func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompositeCodexModelsReusesExistingManifestSelection(t *testing.T) {
|
||||
handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
|
||||
|
||||
recorder := performCodexModelsRequestForPlatform(t, handler, groupID, service.PlatformComposite)
|
||||
|
||||
if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
|
||||
t.Fatalf("upstream account calls: got %v, want %v", got, want)
|
||||
}
|
||||
if recorder.Code != http.StatusOK {
|
||||
t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) {
|
||||
retryableStatuses := []int{
|
||||
http.StatusTooManyRequests,
|
||||
@@ -287,13 +300,17 @@ func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount
|
||||
}
|
||||
|
||||
func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder {
|
||||
return performCodexModelsRequestForPlatform(t, handler, groupID, service.PlatformOpenAI)
|
||||
}
|
||||
|
||||
func performCodexModelsRequestForPlatform(t *testing.T, handler *OpenAIGatewayHandler, groupID int64, platform string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models?client_version=0.144.0", nil)
|
||||
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
|
||||
GroupID: &groupID,
|
||||
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI},
|
||||
Group: &service.Group{ID: groupID, Platform: platform},
|
||||
})
|
||||
|
||||
handler.CodexModels(c)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
coderws "github.com/coder/websocket"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/tidwall/gjson"
|
||||
"github.com/tidwall/sjson"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
@@ -29,7 +30,7 @@ func (h *OpenAIGatewayHandler) Live(c *gin.Context) {
|
||||
h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found")
|
||||
return
|
||||
}
|
||||
if apiKey.Group == nil || apiKey.Group.Platform != service.PlatformOpenAI {
|
||||
if apiKey.Group == nil || (apiKey.Group.Platform != service.PlatformOpenAI && apiKey.Group.Platform != service.PlatformComposite) {
|
||||
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Live is not supported for this platform")
|
||||
return
|
||||
}
|
||||
@@ -43,6 +44,18 @@ func (h *OpenAIGatewayHandler) Live(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
model := strings.TrimSpace(gjson.GetBytes(request.Session, "model").String())
|
||||
if !compositeTargetPlatformAllowed(c, apiKey, model, service.PlatformOpenAI) {
|
||||
h.errorResponse(c, http.StatusNotFound, "not_found_error", "Live only supports OpenAI models for Composite groups")
|
||||
return
|
||||
}
|
||||
if upstreamModel, ok := service.ResolvedUpstreamModelFromContext(c.Request.Context()); ok && upstreamModel != model {
|
||||
rewrittenSession, rewriteErr := sjson.SetBytes(request.Session, "model", upstreamModel)
|
||||
if rewriteErr != nil {
|
||||
h.errorResponse(c, http.StatusInternalServerError, "api_error", "Failed to apply Composite model route")
|
||||
return
|
||||
}
|
||||
request.Session = rewrittenSession
|
||||
}
|
||||
reqLog := requestLogger(
|
||||
c,
|
||||
"handler.openai_gateway.live",
|
||||
@@ -231,6 +244,6 @@ func (h *OpenAIGatewayHandler) LiveSideband(c *gin.Context) {
|
||||
func liveEnabledForAPIKey(apiKey *service.APIKey) bool {
|
||||
return apiKey != nil &&
|
||||
apiKey.Group != nil &&
|
||||
apiKey.Group.Platform == service.PlatformOpenAI &&
|
||||
(apiKey.Group.Platform == service.PlatformOpenAI || apiKey.Group.Platform == service.PlatformComposite) &&
|
||||
apiKey.Group.AllowLive
|
||||
}
|
||||
|
||||
@@ -87,6 +87,9 @@ func TestLiveEnabledForAPIKey(t *testing.T) {
|
||||
require.True(t, liveEnabledForAPIKey(&service.APIKey{
|
||||
Group: &service.Group{Platform: service.PlatformOpenAI, AllowLive: true},
|
||||
}))
|
||||
require.True(t, liveEnabledForAPIKey(&service.APIKey{
|
||||
Group: &service.Group{Platform: service.PlatformComposite, AllowLive: true},
|
||||
}))
|
||||
}
|
||||
|
||||
func TestLiveAttestationErrorIsExplicit(t *testing.T) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -17,15 +17,37 @@ const (
|
||||
|
||||
type toolOutputMediaByCallID map[string][]ChatContentPart
|
||||
|
||||
// ResponsesToChatOptions carries optional hooks for
|
||||
// ResponsesToChatCompletionsRequestWithOptions. All fields are optional; a nil
|
||||
// *ResponsesToChatOptions behaves exactly like ResponsesToChatCompletionsRequest.
|
||||
type ResponsesToChatOptions struct {
|
||||
// ReasoningContentByID looks up the cached reasoning text for a reasoning
|
||||
// item id. Codex histories may carry reasoning items with no plaintext
|
||||
// summary (empty summary + opaque encrypted_content, e.g. after remote
|
||||
// compaction); DeepSeek's thinking mode rejects such histories with 400
|
||||
// "The `reasoning_content` in the thinking mode must be passed back to the
|
||||
// API". The gateway caches the reasoning text it streamed under the item
|
||||
// id, so the lookup restores the reasoning_content the client can no
|
||||
// longer provide. Return "" on a miss. A nil lookup keeps the original
|
||||
// behavior.
|
||||
ReasoningContentByID func(itemID string) string
|
||||
}
|
||||
|
||||
// ResponsesToChatCompletionsRequest converts a Responses API request into a
|
||||
// Chat Completions request for upstreams that only implement
|
||||
// /v1/chat/completions.
|
||||
func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsRequest, error) {
|
||||
return ResponsesToChatCompletionsRequestWithOptions(req, nil)
|
||||
}
|
||||
|
||||
// ResponsesToChatCompletionsRequestWithOptions is ResponsesToChatCompletionsRequest
|
||||
// with optional hooks (see ResponsesToChatOptions).
|
||||
func ResponsesToChatCompletionsRequestWithOptions(req *ResponsesRequest, opts *ResponsesToChatOptions) (*ChatCompletionsRequest, error) {
|
||||
if req == nil {
|
||||
return nil, fmt.Errorf("responses request is nil")
|
||||
}
|
||||
|
||||
messages, err := responsesInputToChatMessages(req.Instructions, req.Input)
|
||||
messages, err := responsesInputToChatMessagesWithOptions(req.Instructions, req.Input, opts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -202,6 +224,12 @@ func HasToolSearchTool(tools []ResponsesTool) bool {
|
||||
// scattered across per-item cases, and makes unknown future codex item types
|
||||
// fail safe instead of leaking into the upstream request.
|
||||
func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage) ([]ChatMessage, error) {
|
||||
return responsesInputToChatMessagesWithOptions(instructions, inputRaw, nil)
|
||||
}
|
||||
|
||||
// responsesInputToChatMessagesWithOptions is responsesInputToChatMessages with
|
||||
// optional hooks (see ResponsesToChatOptions).
|
||||
func responsesInputToChatMessagesWithOptions(instructions string, inputRaw json.RawMessage, opts *ResponsesToChatOptions) ([]ChatMessage, error) {
|
||||
var messages []ChatMessage
|
||||
if strings.TrimSpace(instructions) != "" {
|
||||
content, _ := json.Marshal(instructions)
|
||||
@@ -226,7 +254,7 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
|
||||
return nil, fmt.Errorf("parse responses input: %w", err)
|
||||
}
|
||||
|
||||
built, mediaByCallID, err := buildChatMessagesFromItems(messages, rawItems)
|
||||
built, mediaByCallID, err := buildChatMessagesFromItems(messages, rawItems, opts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -235,7 +263,7 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
|
||||
|
||||
// buildChatMessagesFromItems walks the Responses input items and appends the
|
||||
// corresponding Chat messages.
|
||||
func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessage) ([]ChatMessage, toolOutputMediaByCallID, error) {
|
||||
func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessage, opts *ResponsesToChatOptions) ([]ChatMessage, toolOutputMediaByCallID, error) {
|
||||
// pendingReasoning holds the reasoning text from a reasoning item until the
|
||||
// assistant message it belongs to is emitted. DeepSeek's thinking mode
|
||||
// requires the reasoning_content that produced a tool call to be passed back
|
||||
@@ -243,8 +271,22 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
// across an assistant message (so a following tool call in the same turn
|
||||
// still receives it); any other role ends the thinking span.
|
||||
var pendingReasoning string
|
||||
// lastTurnReasoning is the most recent reasoning text of the current turn,
|
||||
// surviving tool outputs. DeepSeek emits reasoning only once per turn, so
|
||||
// chained tool calls (reasoning → call A → output A → call B) leave call B's
|
||||
// assistant message without reasoning_content and DeepSeek 400s the history;
|
||||
// replaying the turn's reasoning on B's message satisfies the contract. Only
|
||||
// a user-side item ends the turn and clears it.
|
||||
var lastTurnReasoning string
|
||||
mediaByCallID := make(toolOutputMediaByCallID)
|
||||
|
||||
reasoningForAssistant := func() string {
|
||||
if pendingReasoning != "" {
|
||||
return pendingReasoning
|
||||
}
|
||||
return lastTurnReasoning
|
||||
}
|
||||
|
||||
for _, raw := range rawItems {
|
||||
raw = bytesTrimSpace(raw)
|
||||
if len(raw) == 0 || string(raw) == "null" {
|
||||
@@ -258,6 +300,7 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
content, _ := json.Marshal(text)
|
||||
messages = append(messages, ChatMessage{Role: "user", Content: content})
|
||||
pendingReasoning = ""
|
||||
lastTurnReasoning = ""
|
||||
continue
|
||||
}
|
||||
return nil, nil, fmt.Errorf("parse responses input item: %w", err)
|
||||
@@ -269,6 +312,18 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
case "reasoning":
|
||||
if txt := extractResponsesReasoningText(item); txt != "" {
|
||||
pendingReasoning = txt
|
||||
} else if opts != nil && opts.ReasoningContentByID != nil {
|
||||
// No plaintext summary (encrypted-only reasoning, e.g. after codex
|
||||
// remote compaction): fall back to the gateway-side cache keyed
|
||||
// by the reasoning item id, which always round-trips in history.
|
||||
if id := rawString(item["id"]); id != "" {
|
||||
if cached := opts.ReasoningContentByID(id); cached != "" {
|
||||
pendingReasoning = cached
|
||||
}
|
||||
}
|
||||
}
|
||||
if pendingReasoning != "" {
|
||||
lastTurnReasoning = pendingReasoning
|
||||
}
|
||||
continue
|
||||
case "function_call":
|
||||
@@ -290,7 +345,7 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
Arguments: arguments,
|
||||
},
|
||||
}
|
||||
messages = appendAssistantToolCall(messages, toolCall, pendingReasoning)
|
||||
messages = appendAssistantToolCall(messages, toolCall, reasoningForAssistant())
|
||||
pendingReasoning = ""
|
||||
continue
|
||||
case "tool_search_call":
|
||||
@@ -311,7 +366,7 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
Arguments: arguments,
|
||||
},
|
||||
}
|
||||
messages = appendAssistantToolCall(messages, toolCall, pendingReasoning)
|
||||
messages = appendAssistantToolCall(messages, toolCall, reasoningForAssistant())
|
||||
pendingReasoning = ""
|
||||
continue
|
||||
case "custom_tool_call":
|
||||
@@ -327,7 +382,7 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
Arguments: string(arguments),
|
||||
},
|
||||
}
|
||||
messages = appendAssistantToolCall(messages, toolCall, pendingReasoning)
|
||||
messages = appendAssistantToolCall(messages, toolCall, reasoningForAssistant())
|
||||
pendingReasoning = ""
|
||||
continue
|
||||
case "function_call_output", "custom_tool_call_output", "tool_search_output":
|
||||
@@ -359,6 +414,7 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
content, _ := json.Marshal(rawString(item["text"]))
|
||||
messages = append(messages, ChatMessage{Role: "user", Content: content})
|
||||
pendingReasoning = ""
|
||||
lastTurnReasoning = ""
|
||||
continue
|
||||
case "input_image":
|
||||
content, err := chatContentFromSingleResponsesPart(itemType, item)
|
||||
@@ -367,6 +423,7 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
}
|
||||
messages = append(messages, ChatMessage{Role: "user", Content: content})
|
||||
pendingReasoning = ""
|
||||
lastTurnReasoning = ""
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -391,11 +448,22 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
messages = append(messages, ChatMessage{Role: role, Content: chatContent})
|
||||
// Reasoning only survives across an assistant text message.
|
||||
if role != "assistant" {
|
||||
msg := ChatMessage{Role: role, Content: chatContent}
|
||||
// DeepSeek thinking mode requires the reasoning_content from a prior
|
||||
// reasoning-only / plain-text assistant turn to be passed back on its
|
||||
// assistant message; dropping it yields 400 "The `reasoning_content` in
|
||||
// the thinking mode must be passed back to the API" on the next turn.
|
||||
// A following function_call in the same turn still receives it because
|
||||
// appendAssistantToolCall merges into this message and only fills
|
||||
// ReasoningContent when it is still empty.
|
||||
if role == "assistant" {
|
||||
msg.ReasoningContent = reasoningForAssistant()
|
||||
pendingReasoning = ""
|
||||
} else {
|
||||
pendingReasoning = ""
|
||||
lastTurnReasoning = ""
|
||||
}
|
||||
messages = append(messages, msg)
|
||||
}
|
||||
|
||||
return messages, mediaByCallID, nil
|
||||
@@ -687,6 +755,26 @@ func extractResponsesReasoningText(item map[string]json.RawMessage) string {
|
||||
return strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
// ExtractResponsesReasoningItem parses a raw Responses input item and, when it
|
||||
// is a reasoning item, returns its id and extractable plaintext (summary
|
||||
// preferred, content fallback). ok is false for non-reasoning items. It exists
|
||||
// for the gateway-side reasoning cache: items with plaintext get (re)cached so
|
||||
// later encrypted-only replicas of the same item id can be restored.
|
||||
func ExtractResponsesReasoningItem(raw json.RawMessage) (id string, text string, ok bool) {
|
||||
raw = bytesTrimSpace(raw)
|
||||
if len(raw) == 0 || string(raw) == "null" {
|
||||
return "", "", false
|
||||
}
|
||||
var item map[string]json.RawMessage
|
||||
if err := json.Unmarshal(raw, &item); err != nil {
|
||||
return "", "", false
|
||||
}
|
||||
if rawString(item["type"]) != "reasoning" {
|
||||
return "", "", false
|
||||
}
|
||||
return rawString(item["id"]), extractResponsesReasoningText(item), true
|
||||
}
|
||||
|
||||
func chatCompletionsBridgeRole(role string) string {
|
||||
trimmed := strings.TrimSpace(role)
|
||||
if trimmed == "" {
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
package apicompat
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// Encrypted-only reasoning items (empty summary + opaque encrypted_content,
|
||||
// e.g. after codex remote compaction) carry no plaintext the bridge can map to
|
||||
// reasoning_content. The gateway-side cache keyed by reasoning item id restores
|
||||
// it; without the restore, DeepSeek thinking mode rejects the history with 400
|
||||
// "The `reasoning_content` in the thinking mode must be passed back to the API".
|
||||
func TestResponsesToChat_ReasoningCacheLookup_RestoresEncryptedOnlyItem(t *testing.T) {
|
||||
req := &ResponsesRequest{
|
||||
Model: "deepseek-reasoner",
|
||||
Input: json.RawMessage(`[
|
||||
{"type":"reasoning","id":"item_enc1","summary":[],"encrypted_content":"opaque"},
|
||||
{"type":"function_call","call_id":"call_1","name":"get_value","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_1","output":"ok"},
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"go on"}]}
|
||||
]`),
|
||||
}
|
||||
|
||||
out, err := ResponsesToChatCompletionsRequestWithOptions(req, &ResponsesToChatOptions{
|
||||
ReasoningContentByID: func(itemID string) string {
|
||||
if itemID == "item_enc1" {
|
||||
return "cached thinking"
|
||||
}
|
||||
return ""
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Messages, 3)
|
||||
require.Equal(t, "assistant", out.Messages[0].Role)
|
||||
require.Equal(t, "cached thinking", out.Messages[0].ReasoningContent)
|
||||
require.Len(t, out.Messages[0].ToolCalls, 1)
|
||||
require.Equal(t, "call_1", out.Messages[0].ToolCalls[0].ID)
|
||||
require.Equal(t, "tool", out.Messages[1].Role)
|
||||
require.Equal(t, "user", out.Messages[2].Role)
|
||||
}
|
||||
|
||||
// A cache miss keeps the original behavior: no reasoning_content, no error.
|
||||
func TestResponsesToChat_ReasoningCacheLookup_MissKeepsOriginalBehavior(t *testing.T) {
|
||||
req := &ResponsesRequest{
|
||||
Model: "deepseek-reasoner",
|
||||
Input: json.RawMessage(`[
|
||||
{"type":"reasoning","id":"item_unknown","summary":[],"encrypted_content":"opaque"},
|
||||
{"type":"function_call","call_id":"call_1","name":"get_value","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_1","output":"ok"},
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"go on"}]}
|
||||
]`),
|
||||
}
|
||||
|
||||
out, err := ResponsesToChatCompletionsRequestWithOptions(req, &ResponsesToChatOptions{
|
||||
ReasoningContentByID: func(string) string { return "" },
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Messages, 3)
|
||||
require.Empty(t, out.Messages[0].ReasoningContent)
|
||||
|
||||
// Nil options (legacy path) behaves identically.
|
||||
legacy, err := ResponsesToChatCompletionsRequest(req)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, out.Messages, legacy.Messages)
|
||||
}
|
||||
|
||||
// Plaintext summary wins and the cache lookup is not consulted.
|
||||
func TestResponsesToChat_ReasoningCacheLookup_PlaintextPreferred(t *testing.T) {
|
||||
req := &ResponsesRequest{
|
||||
Model: "deepseek-reasoner",
|
||||
Input: json.RawMessage(`[
|
||||
{"type":"reasoning","id":"item_plain","summary":[{"type":"summary_text","text":"plain thinking"}]},
|
||||
{"type":"function_call","call_id":"call_1","name":"get_value","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_1","output":"ok"},
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"go on"}]}
|
||||
]`),
|
||||
}
|
||||
|
||||
lookupCalled := false
|
||||
out, err := ResponsesToChatCompletionsRequestWithOptions(req, &ResponsesToChatOptions{
|
||||
ReasoningContentByID: func(string) string {
|
||||
lookupCalled = true
|
||||
return "cached thinking"
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Messages, 3)
|
||||
require.Equal(t, "plain thinking", out.Messages[0].ReasoningContent)
|
||||
require.False(t, lookupCalled, "plaintext summary present → cache lookup must not run")
|
||||
}
|
||||
|
||||
// DeepSeek emits reasoning only once per turn; chained tool calls
|
||||
// (reasoning → call A → output A → call B) have no reasoning item before call
|
||||
// B. The turn's reasoning must be replayed on B's assistant message, otherwise
|
||||
// DeepSeek thinking mode 400s the history ("reasoning_content ... must be
|
||||
// passed back"). Reproduced from a real codex 0.147.0 resume history.
|
||||
func TestResponsesToChat_ChainedToolCallsReplayTurnReasoning(t *testing.T) {
|
||||
req := &ResponsesRequest{
|
||||
Model: "deepseek-reasoner",
|
||||
Input: json.RawMessage(`[
|
||||
{"type":"reasoning","id":"item_r1","summary":[{"type":"summary_text","text":"turn thinking"}]},
|
||||
{"type":"message","role":"assistant","content":[{"type":"output_text","text":"\n\n"}]},
|
||||
{"type":"function_call","call_id":"call_a","name":"exec_command","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_a","output":"ok"},
|
||||
{"type":"function_call","call_id":"call_b","name":"exec_command","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_b","output":"ok"},
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"next"}]},
|
||||
{"type":"reasoning","id":"item_r2","summary":[{"type":"summary_text","text":"second turn"}]},
|
||||
{"type":"function_call","call_id":"call_c","name":"exec_command","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_c","output":"ok"}
|
||||
]`),
|
||||
}
|
||||
|
||||
out, err := ResponsesToChatCompletionsRequest(req)
|
||||
require.NoError(t, err)
|
||||
|
||||
byCallID := map[string]ChatMessage{}
|
||||
for _, m := range out.Messages {
|
||||
for _, tc := range m.ToolCalls {
|
||||
byCallID[tc.ID] = m
|
||||
}
|
||||
}
|
||||
require.Len(t, byCallID, 3)
|
||||
require.Equal(t, "turn thinking", byCallID["call_a"].ReasoningContent)
|
||||
require.Equal(t, "turn thinking", byCallID["call_b"].ReasoningContent,
|
||||
"链式第二个工具调用必须回放本轮 reasoning")
|
||||
require.Equal(t, "second turn", byCallID["call_c"].ReasoningContent,
|
||||
"user 消息后开启新轮次,不得沿用上一轮 reasoning")
|
||||
|
||||
// 每一条 assistant 消息都必须带 reasoning_content(DeepSeek 契约)。
|
||||
for i, m := range out.Messages {
|
||||
if m.Role == "assistant" {
|
||||
require.NotEmpty(t, m.ReasoningContent, "messages[%d] 缺 reasoning_content", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractResponsesReasoningItem(t *testing.T) {
|
||||
id, text, ok := ExtractResponsesReasoningItem(json.RawMessage(
|
||||
`{"type":"reasoning","id":"item_a","summary":[{"type":"summary_text","text":"think"}]}`))
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "item_a", id)
|
||||
require.Equal(t, "think", text)
|
||||
|
||||
// Encrypted-only item: ok with id but empty text.
|
||||
id, text, ok = ExtractResponsesReasoningItem(json.RawMessage(
|
||||
`{"type":"reasoning","id":"item_b","summary":[],"encrypted_content":"opaque"}`))
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "item_b", id)
|
||||
require.Empty(t, text)
|
||||
|
||||
// Non-reasoning items are skipped.
|
||||
_, _, ok = ExtractResponsesReasoningItem(json.RawMessage(
|
||||
`{"type":"message","role":"user","content":[{"type":"input_text","text":"hi"}]}`))
|
||||
require.False(t, ok)
|
||||
|
||||
_, _, ok = ExtractResponsesReasoningItem(json.RawMessage(`"bare string"`))
|
||||
require.False(t, ok)
|
||||
}
|
||||
@@ -131,6 +131,48 @@ func (c *gatewayCache) ReleaseGrokVideoBilled(ctx context.Context, key string) e
|
||||
var _ service.CyberSessionBlockStore = (*gatewayCache)(nil)
|
||||
var _ service.LiveCallStore = (*gatewayCache)(nil)
|
||||
|
||||
const reasoningContentPrefix = "reasoning_content:"
|
||||
|
||||
// reasoningContentDefaultTTL 是 reasoning 缓存的默认过期时间。Codex 会话可能
|
||||
// 跨多天恢复,取 7 天;调用方传入非正 TTL 时兜底。
|
||||
const reasoningContentDefaultTTL = 7 * 24 * time.Hour
|
||||
|
||||
// SetReasoningContent 按 reasoning item id 缓存 reasoning 全文。
|
||||
// itemID 或 content 为空时直接返回 nil(无可缓存内容,属正常情况而非错误)。
|
||||
func (c *gatewayCache) SetReasoningContent(ctx context.Context, itemID string, content string, ttl time.Duration) error {
|
||||
if c == nil || c.rdb == nil {
|
||||
return errors.New("gateway cache unavailable")
|
||||
}
|
||||
itemID = strings.TrimSpace(itemID)
|
||||
if itemID == "" || content == "" {
|
||||
return nil
|
||||
}
|
||||
if ttl <= 0 {
|
||||
ttl = reasoningContentDefaultTTL
|
||||
}
|
||||
return c.rdb.Set(ctx, reasoningContentPrefix+itemID, content, ttl).Err()
|
||||
}
|
||||
|
||||
// GetReasoningContent 返回缓存的 reasoning 全文;未命中返回
|
||||
// service.ErrReasoningContentNotFound。
|
||||
func (c *gatewayCache) GetReasoningContent(ctx context.Context, itemID string) (string, error) {
|
||||
if c == nil || c.rdb == nil {
|
||||
return "", errors.New("gateway cache unavailable")
|
||||
}
|
||||
itemID = strings.TrimSpace(itemID)
|
||||
if itemID == "" {
|
||||
return "", service.ErrReasoningContentNotFound
|
||||
}
|
||||
val, err := c.rdb.Get(ctx, reasoningContentPrefix+itemID).Result()
|
||||
if err != nil {
|
||||
if errors.Is(err, redis.Nil) {
|
||||
return "", service.ErrReasoningContentNotFound
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
return val, nil
|
||||
}
|
||||
|
||||
const cyberSessionBlockPrefix = "cyber_session_block:"
|
||||
|
||||
// SetCyberSessionBlocked 把被 cyber_policy 命中的会话写入屏蔽表(TTL 自动过期)。
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestGatewayCacheReasoningContent(t *testing.T) {
|
||||
mr := miniredis.RunT(t)
|
||||
client := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
||||
cache := NewGatewayCache(client)
|
||||
ctx := context.Background()
|
||||
|
||||
// 未命中返回哨兵错误,区别于真实读取失败。
|
||||
_, err := cache.GetReasoningContent(ctx, "item_missing")
|
||||
require.ErrorIs(t, err, service.ErrReasoningContentNotFound)
|
||||
|
||||
// 写入后可读回。
|
||||
require.NoError(t, cache.SetReasoningContent(ctx, "item_abc", "think hard", time.Minute))
|
||||
got, err := cache.GetReasoningContent(ctx, "item_abc")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "think hard", got)
|
||||
|
||||
// ttl<=0 时兜底为默认 7 天。
|
||||
require.NoError(t, cache.SetReasoningContent(ctx, "item_ttl", "x", 0))
|
||||
ttl := mr.TTL(reasoningContentPrefix + "item_ttl")
|
||||
require.Greater(t, ttl, 6*24*time.Hour)
|
||||
require.LessOrEqual(t, ttl, reasoningContentDefaultTTL)
|
||||
|
||||
// 空 itemID / 空 content 是 no-op(无可缓存内容不算错误)。
|
||||
require.NoError(t, cache.SetReasoningContent(ctx, "", "x", time.Minute))
|
||||
require.NoError(t, cache.SetReasoningContent(ctx, "item_empty", "", time.Minute))
|
||||
_, err = cache.GetReasoningContent(ctx, "item_empty")
|
||||
require.ErrorIs(t, err, service.ErrReasoningContentNotFound)
|
||||
}
|
||||
@@ -133,6 +133,78 @@ func TestCompositeTargetPlatformMiddlewareUsesExplicitRouteAndRewritesBody(t *te
|
||||
require.Equal(t, http.StatusNoContent, w.Code)
|
||||
}
|
||||
|
||||
func TestCompositeTargetPlatformMiddlewareRewritesNestedLiveModel(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
resolver := service.NewCompositeRouteResolver(compositeRouteRepoStub{
|
||||
routes: []service.CompositeModelRoute{
|
||||
{
|
||||
ID: 1,
|
||||
GroupID: 1,
|
||||
PublicModel: "live-alias",
|
||||
MatchType: service.CompositeRouteMatchExact,
|
||||
TargetPlatform: service.PlatformOpenAI,
|
||||
UpstreamModel: "gpt-live",
|
||||
Endpoint: service.CompositeRouteEndpointAny,
|
||||
Priority: 100,
|
||||
Enabled: true,
|
||||
},
|
||||
},
|
||||
})
|
||||
router.Use(gin.HandlerFunc(servermiddleware.APIKeyAuthMiddleware(func(c *gin.Context) {
|
||||
groupID := int64(1)
|
||||
c.Set(string(servermiddleware.ContextKeyAPIKey), &service.APIKey{
|
||||
GroupID: &groupID,
|
||||
Group: &service.Group{ID: groupID, Platform: service.PlatformComposite},
|
||||
})
|
||||
c.Next()
|
||||
})))
|
||||
router.Use(compositeTargetPlatformMiddleware(resolver))
|
||||
router.POST("/backend-api/codex/realtime/calls", func(c *gin.Context) {
|
||||
platform, ok := service.ResolvedTargetPlatformFromContext(c.Request.Context())
|
||||
require.True(t, ok)
|
||||
require.Equal(t, service.PlatformOpenAI, platform)
|
||||
|
||||
body, err := io.ReadAll(c.Request.Body)
|
||||
require.NoError(t, err)
|
||||
require.JSONEq(t, `{"session":{"model":"gpt-live"},"sdp":"v=0"}`, string(body))
|
||||
c.Status(http.StatusNoContent)
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(
|
||||
http.MethodPost,
|
||||
"/backend-api/codex/realtime/calls",
|
||||
strings.NewReader(`{"session":{"model":"live-alias"},"sdp":"v=0"}`),
|
||||
)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
require.Equal(t, http.StatusNoContent, w.Code)
|
||||
}
|
||||
|
||||
func TestCompositeRequestModelFromMultipartLiveSession(t *testing.T) {
|
||||
var body bytes.Buffer
|
||||
writer := multipart.NewWriter(&body)
|
||||
require.NoError(t, writer.WriteField("sdp", "v=0"))
|
||||
require.NoError(t, writer.WriteField("session", `{"model":"live-alias"}`))
|
||||
require.NoError(t, writer.Close())
|
||||
|
||||
require.Equal(t, "live-alias", compositeRequestModelFromBody(writer.FormDataContentType(), body.Bytes()))
|
||||
}
|
||||
|
||||
func TestCompositeCodexControlPathsUseResponsesRoutes(t *testing.T) {
|
||||
for _, path := range []string{
|
||||
"/v1/alpha/search",
|
||||
"/backend-api/codex/alpha/search",
|
||||
"/v1/live",
|
||||
"/backend-api/codex/realtime/calls",
|
||||
} {
|
||||
require.Equal(t, service.CompositeRouteEndpointResponses, compositeRouteEndpointForPath(path), "path=%s", path)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompositeTargetPlatformMiddlewareUsesExplicitRouteForMultipartImages(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
|
||||
@@ -55,9 +55,6 @@ func RegisterGatewayRoutes(
|
||||
return false
|
||||
}
|
||||
}
|
||||
isOpenAIGatewayPlatform := func(c *gin.Context) bool {
|
||||
return getGroupPlatform(c) == service.PlatformOpenAI
|
||||
}
|
||||
countTokensHandler := func(c *gin.Context) {
|
||||
switch getGroupPlatform(c) {
|
||||
case service.PlatformOpenAI, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek:
|
||||
@@ -69,9 +66,12 @@ func RegisterGatewayRoutes(
|
||||
}
|
||||
}
|
||||
modelsHandler := func(c *gin.Context) {
|
||||
if isOpenAIGatewayPlatform(c) && c.Query("client_version") != "" {
|
||||
h.OpenAIGateway.CodexModels(c)
|
||||
return
|
||||
if c.Query("client_version") != "" {
|
||||
switch getGroupPlatform(c) {
|
||||
case service.PlatformOpenAI, service.PlatformComposite:
|
||||
h.OpenAIGateway.CodexModels(c)
|
||||
return
|
||||
}
|
||||
}
|
||||
h.Gateway.Models(c)
|
||||
}
|
||||
@@ -170,6 +170,10 @@ func RegisterGatewayRoutes(
|
||||
})
|
||||
return
|
||||
}
|
||||
if service.IsOpenAIResponsesInputTokensRequestPath(c) && isOpenAIResponsesCompatibleGatewayPlatform(c) {
|
||||
h.OpenAIGateway.ResponsesInputTokens(c)
|
||||
return
|
||||
}
|
||||
next(c)
|
||||
}
|
||||
}
|
||||
@@ -551,8 +555,10 @@ func compositeTargetPlatformMiddleware(resolver *service.CompositeRouteResolver)
|
||||
if decision.Matched {
|
||||
c.Request = c.Request.WithContext(service.WithCompositeRouteDecision(c.Request.Context(), decision))
|
||||
if upstreamModel := strings.TrimSpace(decision.UpstreamModel); upstreamModel != "" && upstreamModel != model && gjson.ValidBytes(body) {
|
||||
if rewritten, rewriteErr := sjson.SetBytes(body, "model", upstreamModel); rewriteErr == nil {
|
||||
body = rewritten
|
||||
if _, modelPath := compositeJSONRequestModel(body); modelPath != "" {
|
||||
if rewritten, rewriteErr := sjson.SetBytes(body, modelPath, upstreamModel); rewriteErr == nil {
|
||||
body = rewritten
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -563,12 +569,25 @@ func compositeTargetPlatformMiddleware(resolver *service.CompositeRouteResolver)
|
||||
}
|
||||
|
||||
func compositeRequestModelFromBody(contentType string, body []byte) string {
|
||||
if model := strings.TrimSpace(gjson.GetBytes(body, "model").String()); model != "" {
|
||||
if model, _ := compositeJSONRequestModel(body); model != "" {
|
||||
return model
|
||||
}
|
||||
return compositeMultipartModelFromBody(contentType, body)
|
||||
}
|
||||
|
||||
func compositeJSONRequestModel(body []byte) (string, string) {
|
||||
for _, path := range []string{"model", "session.model"} {
|
||||
model := gjson.GetBytes(body, path)
|
||||
if model.Type != gjson.String {
|
||||
continue
|
||||
}
|
||||
if value := strings.TrimSpace(model.String()); value != "" {
|
||||
return value, path
|
||||
}
|
||||
}
|
||||
return "", ""
|
||||
}
|
||||
|
||||
func compositeMultipartModelFromBody(contentType string, body []byte) string {
|
||||
mediaType, params, err := mime.ParseMediaType(strings.TrimSpace(contentType))
|
||||
if err != nil || !strings.EqualFold(mediaType, "multipart/form-data") {
|
||||
@@ -587,14 +606,22 @@ func compositeMultipartModelFromBody(contentType string, body []byte) string {
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
if part.FormName() != "model" || part.FileName() != "" {
|
||||
fieldName := part.FormName()
|
||||
if part.FileName() != "" || (fieldName != "model" && fieldName != "session") {
|
||||
continue
|
||||
}
|
||||
data, err := io.ReadAll(part)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(string(data))
|
||||
switch fieldName {
|
||||
case "model":
|
||||
return strings.TrimSpace(string(data))
|
||||
case "session":
|
||||
if model, _ := compositeJSONRequestModel(data); model != "" {
|
||||
return model
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -670,7 +697,10 @@ func compositeRouteEndpointForPath(path string) string {
|
||||
return service.CompositeRouteEndpointCountTokens
|
||||
case strings.Contains(path, "/messages"):
|
||||
return service.CompositeRouteEndpointMessages
|
||||
case strings.Contains(path, "/responses"):
|
||||
case strings.Contains(path, "/responses"),
|
||||
strings.Contains(path, "/alpha/search"),
|
||||
strings.Contains(path, "/realtime/calls"),
|
||||
strings.HasSuffix(strings.TrimRight(path, "/"), "/live"):
|
||||
return service.CompositeRouteEndpointResponses
|
||||
case strings.Contains(path, "/chat/completions"):
|
||||
return service.CompositeRouteEndpointChatCompletions
|
||||
|
||||
@@ -94,7 +94,7 @@ func TestGatewayRoutesOpenAIAlphaSearchPathsAreRegistered(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayRoutesAlphaSearchRejectsNonOpenAIGroup(t *testing.T) {
|
||||
func TestGatewayRoutesAlphaSearchRejectsUnsupportedGroup(t *testing.T) {
|
||||
router := newGatewayRoutesTestRouter(service.PlatformGrok)
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
@@ -103,7 +103,7 @@ func TestGatewayRoutesAlphaSearchRejectsNonOpenAIGroup(t *testing.T) {
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
require.Equal(t, http.StatusNotFound, w.Code)
|
||||
require.Contains(t, w.Body.String(), "only available for OpenAI groups")
|
||||
require.Contains(t, w.Body.String(), "only available for OpenAI and Composite groups")
|
||||
}
|
||||
|
||||
func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) {
|
||||
|
||||
@@ -512,7 +512,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn
|
||||
ReasoningEffortMappings: reasoningEffortMappings,
|
||||
}
|
||||
sanitizeGroupMessagesDispatchFields(group)
|
||||
if group.Platform != PlatformOpenAI {
|
||||
if group.Platform != PlatformOpenAI && group.Platform != PlatformComposite {
|
||||
group.AllowLive = false
|
||||
}
|
||||
sanitizeGroupReasoningEffortPolicy(group)
|
||||
@@ -892,7 +892,7 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd
|
||||
group.ReasoningEffortMappings = reasoningEffortMappings
|
||||
}
|
||||
sanitizeGroupMessagesDispatchFields(group)
|
||||
if group.Platform != PlatformOpenAI {
|
||||
if group.Platform != PlatformOpenAI && group.Platform != PlatformComposite {
|
||||
group.AllowLive = false
|
||||
}
|
||||
sanitizeGroupReasoningEffortPolicy(group)
|
||||
|
||||
@@ -1029,6 +1029,44 @@ func TestAdminService_CreateGroup_ClearsMessagesDispatchFieldsForNonOpenAIPlatfo
|
||||
require.Equal(t, OpenAIMessagesDispatchModelConfig{}, repo.created.MessagesDispatchModelConfig)
|
||||
}
|
||||
|
||||
func TestAdminService_CreateCompositeGroupPreservesLive(t *testing.T) {
|
||||
repo := &groupRepoStubForAdmin{}
|
||||
svc := &adminServiceImpl{groupRepo: repo}
|
||||
|
||||
group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{
|
||||
Name: "composite-group",
|
||||
Platform: PlatformComposite,
|
||||
RateMultiplier: 1.0,
|
||||
AllowLive: true,
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, group)
|
||||
require.NotNil(t, repo.created)
|
||||
require.True(t, repo.created.AllowLive)
|
||||
}
|
||||
|
||||
func TestAdminService_UpdateCompositeGroupPreservesLive(t *testing.T) {
|
||||
existingGroup := &Group{
|
||||
ID: 1,
|
||||
Name: "composite-group",
|
||||
Platform: PlatformComposite,
|
||||
Status: StatusActive,
|
||||
}
|
||||
repo := &groupRepoStubForAdmin{getByID: existingGroup}
|
||||
svc := &adminServiceImpl{groupRepo: repo}
|
||||
allowLive := true
|
||||
|
||||
group, err := svc.UpdateGroup(context.Background(), existingGroup.ID, &UpdateGroupInput{
|
||||
AllowLive: &allowLive,
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, group)
|
||||
require.NotNil(t, repo.updated)
|
||||
require.True(t, repo.updated.AllowLive)
|
||||
}
|
||||
|
||||
func TestAdminService_UpdateGroup_ClearsMessagesDispatchFieldsWhenPlatformChangesAwayFromOpenAI(t *testing.T) {
|
||||
existingGroup := &Group{
|
||||
ID: 1,
|
||||
|
||||
@@ -158,6 +158,13 @@ func (s *stickyGatewayCacheHotpathStub) ReleaseGrokVideoBilled(_ context.Context
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *stickyGatewayCacheHotpathStub) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (s *stickyGatewayCacheHotpathStub) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
func (s *modelsListAccountRepoStub) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]Account, error) {
|
||||
s.listByGroupCalls.Add(1)
|
||||
if s.err != nil {
|
||||
|
||||
@@ -291,6 +291,13 @@ func (m *mockGatewayCacheForPlatform) ReleaseGrokVideoBilled(_ context.Context,
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockGatewayCacheForPlatform) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (m *mockGatewayCacheForPlatform) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
type mockGroupRepoForGateway struct {
|
||||
groups map[int64]*Group
|
||||
getByIDCalls int
|
||||
|
||||
@@ -452,6 +452,10 @@ var allowedHeaders = map[string]bool{
|
||||
// cache implementation (e.g. redis.Nil), mirroring ErrRefreshTokenNotFound.
|
||||
var ErrStickySessionNotFound = errors.New("sticky session not found")
|
||||
|
||||
// ErrReasoningContentNotFound is returned by GatewayCache.GetReasoningContent
|
||||
// when no cached reasoning content exists for the reasoning item ID.
|
||||
var ErrReasoningContentNotFound = errors.New("reasoning content not found")
|
||||
|
||||
// GatewayCache 定义网关服务的缓存操作接口。
|
||||
// 提供粘性会话(Sticky Session)的存储、查询、刷新和删除功能。
|
||||
//
|
||||
@@ -486,6 +490,16 @@ type GatewayCache interface {
|
||||
ClaimGrokVideoBilled(ctx context.Context, key string, ttl time.Duration) (bool, error)
|
||||
// ReleaseGrokVideoBilled clears a claim so a failed RecordUsage can retry billing.
|
||||
ReleaseGrokVideoBilled(ctx context.Context, key string) error
|
||||
|
||||
// Reasoning content cache (Responses→Chat Completions 桥接)。
|
||||
// SetReasoningContent 按 reasoning item id 缓存 reasoning 全文,供后续请求
|
||||
// 在客户端不回传明文 summary 时回注 reasoning_content(DeepSeek thinking
|
||||
// mode 要求回传,否则 400)。
|
||||
SetReasoningContent(ctx context.Context, itemID string, content string, ttl time.Duration) error
|
||||
// GetReasoningContent 返回缓存的 reasoning 全文;未命中返回
|
||||
// ErrReasoningContentNotFound,使 service 层无需依赖具体缓存实现即可
|
||||
// 区分"未缓存"与真实读取失败。
|
||||
GetReasoningContent(ctx context.Context, itemID string) (string, error)
|
||||
}
|
||||
|
||||
// derefGroupID safely dereferences *int64 to int64, returning 0 if nil
|
||||
|
||||
@@ -319,6 +319,13 @@ func (m *mockGatewayCacheForGemini) ReleaseGrokVideoBilled(_ context.Context, _
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockGatewayCacheForGemini) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (m *mockGatewayCacheForGemini) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
// TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform 测试 Gemini 单平台选择
|
||||
func TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
@@ -196,6 +196,13 @@ func (c *schedulerTestGatewayCache) ReleaseGrokVideoBilled(_ context.Context, _
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *schedulerTestGatewayCache) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (c *schedulerTestGatewayCache) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
func newSchedulerTestOpenAIWSV2Config() *config.Config {
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.Enabled = true
|
||||
|
||||
@@ -148,6 +148,13 @@ func (c *comboCacheAndStore) ReleaseGrokVideoBilled(_ context.Context, _ string)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *comboCacheAndStore) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (c *comboCacheAndStore) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
func (c *comboCacheAndStore) SetCyberSessionBlocked(ctx context.Context, key string, ttl time.Duration) error {
|
||||
return c.store.SetCyberSessionBlocked(ctx, key, ttl)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -52,7 +53,13 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
|
||||
toolSearch := apicompat.HasToolSearchTool(effectiveTools)
|
||||
namespaceTools := apicompat.NamespaceToolNames(effectiveTools)
|
||||
|
||||
chatReq, err := apicompat.ResponsesToChatCompletionsRequest(&responsesReq)
|
||||
// 自愈回写:历史里带明文 summary 的 reasoning item 刷新进缓存,覆盖 Redis
|
||||
// 被 flush / 跨实例漂移后同 id 的 encrypted-only 副本无法再取明文的情况。
|
||||
s.recacheReasoningItemsFromInput(responsesReq.Input)
|
||||
|
||||
chatReq, err := apicompat.ResponsesToChatCompletionsRequestWithOptions(&responsesReq, &apicompat.ResponsesToChatOptions{
|
||||
ReasoningContentByID: s.reasoningContentByID,
|
||||
})
|
||||
if err != nil {
|
||||
writeOpenAIResponsesFallbackError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
||||
return nil, fmt.Errorf("convert responses to chat completions: %w", err)
|
||||
@@ -136,6 +143,7 @@ func (s *OpenAIGatewayService) bufferChatCompletionsAsResponses(
|
||||
return nil, err
|
||||
}
|
||||
responsesResp := apicompat.ChatCompletionsResponseToResponses(ccResp, originalModel, customTools, toolSearch, namespaceTools)
|
||||
s.cacheReasoningItemsFromOutput(responsesResp.Output)
|
||||
|
||||
if s.responseHeaderFilter != nil {
|
||||
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
|
||||
@@ -204,7 +212,9 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses(
|
||||
}
|
||||
|
||||
scan := s.scanCCStream(resp, "openai responses chat fallback", requestID, startTime, func(chunk *apicompat.ChatCompletionsChunk) {
|
||||
writeEvents(apicompat.ChatCompletionsChunkToResponsesEvents(chunk, state))
|
||||
events := apicompat.ChatCompletionsChunkToResponsesEvents(chunk, state)
|
||||
s.cacheReasoningItemsFromEvents(events)
|
||||
writeEvents(events)
|
||||
})
|
||||
|
||||
if scan.Err != nil {
|
||||
@@ -222,7 +232,9 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses(
|
||||
}, fmt.Errorf("stream usage incomplete: %w", scan.Err)
|
||||
}
|
||||
|
||||
writeEvents(apicompat.FinalizeChatCompletionsResponsesStream(state))
|
||||
finalEvents := apicompat.FinalizeChatCompletionsResponsesStream(state)
|
||||
s.cacheReasoningItemsFromEvents(finalEvents)
|
||||
writeEvents(finalEvents)
|
||||
if !clientDisconnected {
|
||||
writeStreamHeaders()
|
||||
if _, err := fmt.Fprint(c.Writer, "data: [DONE]\n\n"); err != nil {
|
||||
@@ -261,3 +273,100 @@ func chatChunkStartsResponsesOutput(chunk *apicompat.ChatCompletionsChunk) bool
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// responsesReasoningCacheTTL 是 reasoning 缓存(按 reasoning item id)的过期时间。
|
||||
// Codex 会话可能跨多天恢复历史,取 7 天。
|
||||
const responsesReasoningCacheTTL = 7 * 24 * time.Hour
|
||||
|
||||
// reasoningContentByID 按 reasoning item id 回查缓存的 reasoning 全文,供
|
||||
// Responses→CC 桥接在客户端不回传明文 summary(encrypted-only reasoning
|
||||
// item)时回注 reasoning_content。任何失败都 fail-open 返回 ""(维持桥接原
|
||||
// 行为),因为缓存只是优化而非正确性前提。
|
||||
func (s *OpenAIGatewayService) reasoningContentByID(itemID string) string {
|
||||
if s == nil || s.cache == nil {
|
||||
return ""
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
content, err := s.cache.GetReasoningContent(ctx, itemID)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return content
|
||||
}
|
||||
|
||||
// recacheReasoningItemsFromInput 把请求历史里带明文 summary 的 reasoning item
|
||||
// 重新写入缓存(best-effort)。Codex 多数时候会原样回传明文 summary,借机
|
||||
// 刷新 TTL 并自愈 Redis 被 flush / 跨实例漂移造成的缓存缺失。
|
||||
func (s *OpenAIGatewayService) recacheReasoningItemsFromInput(inputRaw json.RawMessage) {
|
||||
if s == nil || s.cache == nil {
|
||||
return
|
||||
}
|
||||
inputRaw = bytes.TrimSpace(inputRaw)
|
||||
if len(inputRaw) == 0 || inputRaw[0] != '[' {
|
||||
return
|
||||
}
|
||||
var items []json.RawMessage
|
||||
if err := json.Unmarshal(inputRaw, &items); err != nil {
|
||||
return
|
||||
}
|
||||
for _, raw := range items {
|
||||
id, text, ok := apicompat.ExtractResponsesReasoningItem(raw)
|
||||
if !ok || id == "" || text == "" {
|
||||
continue
|
||||
}
|
||||
s.setReasoningContent(id, text)
|
||||
}
|
||||
}
|
||||
|
||||
// cacheReasoningItemsFromEvents 从 Responses 流事件里提取完成的 reasoning
|
||||
// item 写入缓存(覆盖一个流中的多个 reasoning item)。
|
||||
func (s *OpenAIGatewayService) cacheReasoningItemsFromEvents(events []apicompat.ResponsesStreamEvent) {
|
||||
for _, event := range events {
|
||||
if event.Type != "response.output_item.done" || event.Item == nil {
|
||||
continue
|
||||
}
|
||||
s.cacheReasoningItem(event.Item)
|
||||
}
|
||||
}
|
||||
|
||||
// cacheReasoningItemsFromOutput 从非流式 Responses 响应的 output 里提取
|
||||
// reasoning item 写入缓存。
|
||||
func (s *OpenAIGatewayService) cacheReasoningItemsFromOutput(output []apicompat.ResponsesOutput) {
|
||||
for i := range output {
|
||||
s.cacheReasoningItem(&output[i])
|
||||
}
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) cacheReasoningItem(item *apicompat.ResponsesOutput) {
|
||||
if item == nil || item.Type != "reasoning" || item.ID == "" {
|
||||
return
|
||||
}
|
||||
var parts []string
|
||||
for _, sum := range item.Summary {
|
||||
if t := strings.TrimSpace(sum.Text); t != "" {
|
||||
parts = append(parts, t)
|
||||
}
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return
|
||||
}
|
||||
s.setReasoningContent(item.ID, strings.Join(parts, "\n"))
|
||||
}
|
||||
|
||||
// setReasoningContent 写入缓存,使用 detached ctx:客户端断连后仍在 drain
|
||||
// 上游流(计费需要),此时的 reasoning 也是后续轮次回注所依赖的,不能随
|
||||
// 请求 ctx 一起取消。失败仅记日志,不影响转发。
|
||||
func (s *OpenAIGatewayService) setReasoningContent(itemID, content string) {
|
||||
if s == nil || s.cache == nil {
|
||||
return
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
if err := s.cache.SetReasoningContent(ctx, itemID, content, responsesReasoningCacheTTL); err != nil {
|
||||
logger.L().Warn("openai responses chat fallback: cache reasoning content failed",
|
||||
zap.Error(err),
|
||||
zap.String("item_id", itemID),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,7 +9,9 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -227,3 +229,143 @@ func forceChatResponsesFallbackAccount() *Account {
|
||||
}
|
||||
return account
|
||||
}
|
||||
|
||||
// reasoningRecordingCache 记录 reasoning 缓存写入、并按需响应回查。
|
||||
type reasoningRecordingCache struct {
|
||||
stubGatewayCache
|
||||
mu sync.Mutex
|
||||
sets map[string]string
|
||||
getResp map[string]string
|
||||
}
|
||||
|
||||
func (c *reasoningRecordingCache) SetReasoningContent(_ context.Context, itemID string, content string, _ time.Duration) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.sets == nil {
|
||||
c.sets = make(map[string]string)
|
||||
}
|
||||
c.sets[itemID] = content
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *reasoningRecordingCache) GetReasoningContent(_ context.Context, itemID string) (string, error) {
|
||||
if v, ok := c.getResp[itemID]; ok {
|
||||
return v, nil
|
||||
}
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
func (c *reasoningRecordingCache) snapshotSets() map[string]string {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
out := make(map[string]string, len(c.sets))
|
||||
for k, v := range c.sets {
|
||||
out[k] = v
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// 流式响应里的 reasoning_content 应按 reasoning item id 写入缓存,供后续轮次
|
||||
// 客户端不回传明文 summary 时回注(DeepSeek thinking mode 400 修复的写入侧)。
|
||||
func TestForwardResponses_ChatFallbackCachesStreamedReasoning(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-reasoner","input":"hello","stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
upstreamBody := strings.Join([]string{
|
||||
`data: {"id":"chatcmpl_rc","object":"chat.completion.chunk","model":"deepseek-reasoner","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_rc","object":"chat.completion.chunk","model":"deepseek-reasoner","choices":[{"index":0,"delta":{"reasoning_content":"think "},"finish_reason":null}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_rc","object":"chat.completion.chunk","model":"deepseek-reasoner","choices":[{"index":0,"delta":{"reasoning_content":"first"},"finish_reason":null}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_rc","object":"chat.completion.chunk","model":"deepseek-reasoner","choices":[{"index":0,"delta":{"content":"answer"},"finish_reason":"stop"}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_rc","object":"chat.completion.chunk","model":"deepseek-reasoner","choices":[],"usage":{"prompt_tokens":4,"completion_tokens":3,"total_tokens":7}}`,
|
||||
"",
|
||||
"data: [DONE]",
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_reasoning_cache_stream"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||
}}
|
||||
cache := &reasoningRecordingCache{}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
cache: cache,
|
||||
}
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, forceChatResponsesFallbackAccount(), body)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
|
||||
sets := cache.snapshotSets()
|
||||
require.Len(t, sets, 1, "应恰好缓存一个 reasoning item")
|
||||
for itemID, content := range sets {
|
||||
require.NotEmpty(t, itemID)
|
||||
require.Equal(t, "think first", content)
|
||||
}
|
||||
}
|
||||
|
||||
// 请求侧:encrypted-only reasoning item(无明文 summary)经缓存回查补回
|
||||
// reasoning_content;带明文 summary 的 item 顺手回写缓存(自愈)。
|
||||
func TestForwardResponses_ChatFallbackRestoresReasoningFromCache(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{
|
||||
"model":"deepseek-reasoner",
|
||||
"stream":false,
|
||||
"input":[
|
||||
{"type":"reasoning","id":"item_plain","summary":[{"type":"summary_text","text":"plain thinking"}]},
|
||||
{"type":"function_call","call_id":"call_0","name":"get_value","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_0","output":"ok"},
|
||||
{"type":"reasoning","id":"item_enc1","summary":[],"encrypted_content":"opaque"},
|
||||
{"type":"function_call","call_id":"call_1","name":"get_value","arguments":"{}"},
|
||||
{"type":"function_call_output","call_id":"call_1","output":"ok"},
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"go on"}]}
|
||||
]
|
||||
}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", 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"}, "x-request-id": []string{"rid_reasoning_cache_restore"}},
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
`{"id":"chatcmpl_restore","object":"chat.completion","model":"deepseek-reasoner","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}`,
|
||||
)),
|
||||
}}
|
||||
cache := &reasoningRecordingCache{
|
||||
getResp: map[string]string{"item_enc1": "cached thinking"},
|
||||
}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
cache: cache,
|
||||
}
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, forceChatResponsesFallbackAccount(), body)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
|
||||
// 明文 summary 的 assistant 工具调用消息:reasoning_content 来自 summary 本身。
|
||||
require.Equal(t, "plain thinking", gjson.GetBytes(upstream.lastBody, "messages.0.reasoning_content").String())
|
||||
require.Equal(t, "call_0", gjson.GetBytes(upstream.lastBody, "messages.0.tool_calls.0.id").String())
|
||||
require.Equal(t, "tool", gjson.GetBytes(upstream.lastBody, "messages.1.role").String())
|
||||
// encrypted-only 的 assistant 工具调用消息:reasoning_content 来自缓存回查。
|
||||
require.Equal(t, "cached thinking", gjson.GetBytes(upstream.lastBody, "messages.2.reasoning_content").String())
|
||||
require.Equal(t, "call_1", gjson.GetBytes(upstream.lastBody, "messages.2.tool_calls.0.id").String())
|
||||
require.Equal(t, "tool", gjson.GetBytes(upstream.lastBody, "messages.3.role").String())
|
||||
|
||||
// 明文 summary 的 item 被回写进缓存(自愈)。
|
||||
require.Equal(t, "plain thinking", cache.snapshotSets()["item_plain"])
|
||||
}
|
||||
|
||||
@@ -712,6 +712,13 @@ func (c *stubGatewayCache) ReleaseGrokVideoBilled(_ context.Context, _ string) e
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *stubGatewayCache) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (c *stubGatewayCache) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
func TestOpenAISelectAccountWithLoadAwareness_FiltersUnschedulable(t *testing.T) {
|
||||
now := time.Now()
|
||||
resetAt := now.Add(10 * time.Minute)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -207,6 +207,13 @@ func (c *openAIWSStateStoreTimeoutProbeCache) ReleaseGrokVideoBilled(_ context.C
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *openAIWSStateStoreTimeoutProbeCache) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (c *openAIWSStateStoreTimeoutProbeCache) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
func TestOpenAIWSStateStore_RedisOpsUseShortTimeout(t *testing.T) {
|
||||
probe := &openAIWSStateStoreTimeoutProbeCache{}
|
||||
store := NewOpenAIWSStateStore(probe)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -118,6 +118,13 @@ func (c StubGatewayCache) ReleaseGrokVideoBilled(_ context.Context, _ string) er
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c StubGatewayCache) SetReasoningContent(_ context.Context, _ string, _ string, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
func (c StubGatewayCache) GetReasoningContent(_ context.Context, _ string) (string, error) {
|
||||
return "", service.ErrReasoningContentNotFound
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// StubSessionLimitCache — service.SessionLimitCache 的空实现
|
||||
// ============================================================
|
||||
|
||||
@@ -53,6 +53,12 @@ route's `upstream_model` before dispatch. For Gemini native paths such as
|
||||
`/v1beta/models/{model}:generateContent`, the gateway resolves `{model}` and
|
||||
the handler forwards the resolved upstream model.
|
||||
|
||||
Codex Alpha Search and Live requests use the `responses` route domain. Live
|
||||
requests resolve the model from `session.model`, including multipart `session`
|
||||
payloads, and apply the configured `upstream_model` before dispatch.
|
||||
Codex model manifest requests reuse the existing OpenAI account selection and
|
||||
failover path within the Composite group.
|
||||
|
||||
## Built-In Detection
|
||||
|
||||
Composite routing detects common public model IDs and provider-prefixed IDs:
|
||||
|
||||
@@ -1581,9 +1581,9 @@
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<!-- OpenAI Live 开关(仅 openai 平台) -->
|
||||
<!-- Codex Live 开关(OpenAI 与 Composite 平台) -->
|
||||
<div
|
||||
v-if="createForm.platform === 'openai'"
|
||||
v-if="supportsLivePlatform(createForm.platform)"
|
||||
class="border-t border-gray-200 dark:border-dark-400 pt-4 mt-4"
|
||||
>
|
||||
<h4 class="text-sm font-medium text-gray-700 dark:text-gray-300 mb-3">
|
||||
@@ -3303,9 +3303,9 @@
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<!-- OpenAI Live 开关(仅 openai 平台) -->
|
||||
<!-- Codex Live 开关(OpenAI 与 Composite 平台) -->
|
||||
<div
|
||||
v-if="editForm.platform === 'openai'"
|
||||
v-if="supportsLivePlatform(editForm.platform)"
|
||||
class="border-t border-gray-200 dark:border-dark-400 pt-4 mt-4"
|
||||
>
|
||||
<h4 class="text-sm font-medium text-gray-700 dark:text-gray-300 mb-3">
|
||||
@@ -4501,6 +4501,9 @@ import {
|
||||
videoModelPriceFamilyRows,
|
||||
} from "./groupsVideoModelPricing";
|
||||
|
||||
const supportsLivePlatform = (platform: string): boolean =>
|
||||
platform === "openai" || platform === "composite";
|
||||
|
||||
const emptyGroupPricing = (): PricingFormEntry => ({
|
||||
models: [],
|
||||
billing_mode: "token",
|
||||
@@ -6666,6 +6669,8 @@ watch(
|
||||
}
|
||||
if (newVal !== "openai") {
|
||||
resetMessagesDispatchFormState(createForm);
|
||||
}
|
||||
if (!supportsLivePlatform(newVal)) {
|
||||
createForm.allow_live = false;
|
||||
}
|
||||
if (!isProfitControlPlatform(newVal)) {
|
||||
@@ -6714,6 +6719,8 @@ watch(
|
||||
}
|
||||
if (newVal !== "openai") {
|
||||
resetMessagesDispatchFormState(editForm);
|
||||
}
|
||||
if (!supportsLivePlatform(newVal)) {
|
||||
editForm.allow_live = false;
|
||||
}
|
||||
if (!isProfitControlPlatform(newVal)) {
|
||||
@@ -6764,9 +6771,11 @@ watch(
|
||||
}
|
||||
if (newVal !== 'openai') {
|
||||
editForm.allow_messages_dispatch = false
|
||||
editForm.allow_live = false
|
||||
editForm.default_mapped_model = ''
|
||||
}
|
||||
if (!supportsLivePlatform(newVal)) {
|
||||
editForm.allow_live = false
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user