Merge remote-tracking branch 'origin/main' into fix/issue-5796-composite-new-platforms

This commit is contained in:
shaw
2026-08-19 14:47:34 +08:00
36 changed files with 1360 additions and 57 deletions
+27 -15
View File
@@ -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 -2
View File
@@ -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"))
}
+6 -2
View File
@@ -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()
+42 -12
View File
@@ -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) {
+2 -2
View File
@@ -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",
+7
View File
@@ -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 的空实现
// ============================================================
+6
View File
@@ -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:
+14 -5
View File
@@ -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
}
}
)