mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:43:23 +08:00
Merge branch 'main' into fix/anthropic
This commit is contained in:
@@ -1307,11 +1307,15 @@ func writeGrokModelsList(c *gin.Context, modelIDs []string) {
|
||||
if grokModelSupportsConfigurableReasoning(modelID) {
|
||||
item.SupportsReasoningEffort = true
|
||||
item.ReasoningEffort = "high"
|
||||
item.ReasoningEfforts = []grokReasoningEffortOption{
|
||||
efforts := []grokReasoningEffortOption{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
}
|
||||
if service.GrokSupportsXHighReasoningEffort(modelID) {
|
||||
efforts = append(efforts, grokReasoningEffortOption{Value: "xhigh", Label: "xHigh"})
|
||||
}
|
||||
item.ReasoningEfforts = efforts
|
||||
}
|
||||
models = append(models, item)
|
||||
}
|
||||
|
||||
@@ -364,9 +364,38 @@ func TestGatewayModels_GeminiGroupFallsBackToGeminiModels(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T) {
|
||||
assertGrokGatewayReasoningEfforts(t, 4409, "grok-4.5", []gatewayReasoningEffortOptionForTest{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
})
|
||||
}
|
||||
|
||||
func TestGatewayModels_Grok46AdvertisesXHighReasoningEffortForGrokBuild(t *testing.T) {
|
||||
xhighEfforts := []gatewayReasoningEffortOptionForTest{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
{Value: "xhigh", Label: "xHigh"},
|
||||
}
|
||||
tests := []struct {
|
||||
groupID int64
|
||||
model string
|
||||
}{
|
||||
{groupID: 4410, model: "grok-4.6"},
|
||||
{groupID: 4411, model: "grok-4.6-latest"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.model, func(t *testing.T) {
|
||||
assertGrokGatewayReasoningEfforts(t, tt.groupID, tt.model, xhighEfforts)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func assertGrokGatewayReasoningEfforts(t *testing.T, groupID int64, modelID string, want []gatewayReasoningEffortOptionForTest) {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
groupID := int64(4409)
|
||||
h := newGatewayModelsHandlerForTest(
|
||||
&gatewayModelsAccountRepoStub{
|
||||
byGroup: map[int64][]service.Account{
|
||||
@@ -375,7 +404,7 @@ func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T)
|
||||
ID: 1,
|
||||
Platform: service.PlatformGrok,
|
||||
Credentials: map[string]any{
|
||||
"model_mapping": map[string]any{"grok-4.5": "grok-4.5"},
|
||||
"model_mapping": map[string]any{modelID: modelID},
|
||||
},
|
||||
},
|
||||
},
|
||||
@@ -397,14 +426,10 @@ func TestGatewayModels_Grok45AdvertisesReasoningEffortForGrokBuild(t *testing.T)
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
require.Len(t, got.Data, 1)
|
||||
model := got.Data[0]
|
||||
require.Equal(t, "grok-4.5", model.ID)
|
||||
require.Equal(t, modelID, model.ID)
|
||||
require.True(t, model.SupportsReasoningEffort)
|
||||
require.Equal(t, "high", model.ReasoningEffort)
|
||||
require.Equal(t, []gatewayReasoningEffortOptionForTest{
|
||||
{Value: "low", Label: "Low"},
|
||||
{Value: "medium", Label: "Medium"},
|
||||
{Value: "high", Label: "High", Default: true},
|
||||
}, model.ReasoningEfforts)
|
||||
require.Equal(t, want, model.ReasoningEfforts)
|
||||
}
|
||||
|
||||
func TestGatewayModels_GeminiGroupFiltersMappedModelsByPlatform(t *testing.T) {
|
||||
|
||||
@@ -2274,6 +2274,13 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
c.Set(securityAuditWSTurnContextKey, turn)
|
||||
service.BeginOpsStreamTurn(c, turn)
|
||||
setCyberTurnBody(turn, payload)
|
||||
// Passthrough ingress intentionally skips BeforeTurn, so enforce only
|
||||
// the connection-level cyber session gate here as well. Native ingress
|
||||
// visits this hook first and gets the same side-effect-free close error;
|
||||
// its original BeforeTurn guard remains as defense in depth.
|
||||
if cyberBlockedThisConn {
|
||||
return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, cyberSessionBlockedClientMsg, nil)
|
||||
}
|
||||
if turn == 1 {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,319 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/Wei-Shaw/sub2api/internal/testutil"
|
||||
coderws "github.com/coder/websocket"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
type openAIWSPassthroughHandlerHarness struct {
|
||||
clientConn *coderws.Conn
|
||||
handlerDone <-chan struct{}
|
||||
moderationRepo *contentModerationHandlerTestRepo
|
||||
gatewayCache service.GatewayCache
|
||||
apiKey *service.APIKey
|
||||
}
|
||||
|
||||
func newOpenAIWSPassthroughHandlerHarness(t *testing.T, upstreamURL string) *openAIWSPassthroughHandlerHarness {
|
||||
t.Helper()
|
||||
gatewayCache := testutil.NewRedisGatewayCache(t)
|
||||
|
||||
settingRepo := &contentModerationHandlerSettingRepo{values: map[string]string{
|
||||
service.SettingKeyRiskControlEnabled: "true",
|
||||
service.SettingKeyCyberSessionBlockEnabled: "true",
|
||||
service.SettingKeyCyberSessionBlockTTLSeconds: "60",
|
||||
}}
|
||||
moderationRepo := &contentModerationHandlerTestRepo{}
|
||||
moderationSvc := service.NewContentModerationService(settingRepo, moderationRepo, nil, nil, nil, nil, nil, nil)
|
||||
settingSvc := service.NewSettingService(settingRepo, nil)
|
||||
|
||||
groupID := int64(4301)
|
||||
account := service.Account{
|
||||
ID: 9951,
|
||||
Name: "openai-ws-passthrough-cyber",
|
||||
Platform: service.PlatformOpenAI,
|
||||
Type: service.AccountTypeAPIKey,
|
||||
Status: service.StatusActive,
|
||||
Schedulable: true,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{"api_key": "sk-test", "base_url": upstreamURL},
|
||||
Extra: map[string]any{
|
||||
"openai_apikey_responses_websockets_v2_enabled": true,
|
||||
"openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
|
||||
},
|
||||
}
|
||||
cfg := &config.Config{}
|
||||
cfg.RunMode = config.RunModeSimple
|
||||
cfg.Default.RateMultiplier = 1
|
||||
cfg.Security.URLAllowlist.Enabled = false
|
||||
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
|
||||
cfg.Gateway.OpenAIWS.Enabled = true
|
||||
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
|
||||
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
|
||||
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
|
||||
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
|
||||
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
|
||||
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
|
||||
cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds = 3
|
||||
|
||||
accountRepo := &openAIWSUsageHandlerAccountRepoStub{account: account}
|
||||
usageRepo := &openAIWSUsageHandlerUsageLogRepoStub{created: make(chan *service.UsageLog, 2)}
|
||||
billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil)
|
||||
gatewaySvc := service.NewOpenAIGatewayService(
|
||||
accountRepo, usageRepo, nil, nil, nil, nil, gatewayCache, cfg, nil, nil,
|
||||
service.NewBillingService(cfg, nil), nil, billingCacheSvc, nil, &service.DeferredService{},
|
||||
nil, nil, nil, nil, nil, settingSvc, nil,
|
||||
)
|
||||
concurrencyCache := &concurrencyCacheMock{
|
||||
acquireUserSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil },
|
||||
acquireAccountSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil },
|
||||
}
|
||||
h := &OpenAIGatewayHandler{
|
||||
gatewayService: gatewaySvc,
|
||||
billingCacheService: billingCacheSvc,
|
||||
apiKeyService: &service.APIKeyService{},
|
||||
contentModerationService: moderationSvc,
|
||||
concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(concurrencyCache), SSEPingFormatNone, time.Second),
|
||||
}
|
||||
|
||||
apiKey := &service.APIKey{
|
||||
ID: 1851,
|
||||
Name: "ws-cyber-key",
|
||||
Key: "sk-handler-cyber-test",
|
||||
GroupID: &groupID,
|
||||
User: &service.User{ID: 1751, Status: service.StatusActive},
|
||||
}
|
||||
handlerDone := make(chan struct{})
|
||||
router := gin.New()
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Set(string(middleware.ContextKeyAPIKey), apiKey)
|
||||
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.User.ID, Concurrency: 1})
|
||||
c.Next()
|
||||
})
|
||||
router.GET("/openai/v1/responses", func(c *gin.Context) {
|
||||
h.ResponsesWebSocket(c)
|
||||
close(handlerDone)
|
||||
})
|
||||
handlerServer := httptest.NewServer(router)
|
||||
t.Cleanup(handlerServer.Close)
|
||||
|
||||
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(handlerServer.URL, "http")+"/openai/v1/responses", nil)
|
||||
cancelDial()
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { _ = clientConn.CloseNow() })
|
||||
|
||||
return &openAIWSPassthroughHandlerHarness{
|
||||
clientConn: clientConn,
|
||||
handlerDone: handlerDone,
|
||||
moderationRepo: moderationRepo,
|
||||
gatewayCache: gatewayCache,
|
||||
apiKey: apiKey,
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesWebSocketV2PassthroughCyberMarkIsConsumedAfterTurn(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
upstreamDone := make(chan struct{})
|
||||
secondUpstreamFrame := make(chan []byte, 1)
|
||||
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
defer close(upstreamDone)
|
||||
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = conn.CloseNow() }()
|
||||
|
||||
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
_, _, err = conn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, err)
|
||||
|
||||
failed := []byte(`{"type":"response.failed","response":{"id":"resp_cyber_handler","model":"gpt-5.1","error":{"code":"cyber_policy","message":"blocked by upstream policy"},"usage":{"input_tokens":11,"output_tokens":3}}}`)
|
||||
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
err = conn.Write(writeCtx, coderws.MessageText, failed)
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second)
|
||||
_, second, err := conn.Read(readCtx)
|
||||
cancelRead()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
secondUpstreamFrame <- append([]byte(nil), second...)
|
||||
|
||||
completed := []byte(`{"type":"response.completed","response":{"id":"resp_cyber_handler_turn_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`)
|
||||
writeCtx, cancelWrite = context.WithTimeout(r.Context(), 3*time.Second)
|
||||
err = conn.Write(writeCtx, coderws.MessageText, completed)
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
}))
|
||||
defer upstreamServer.Close()
|
||||
harness := newOpenAIWSPassthroughHandlerHarness(t, upstreamServer.URL)
|
||||
|
||||
requestPayload := `{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"cyber-session-1","input":"test"}`
|
||||
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err := harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(requestPayload))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, event, err := harness.clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "response.failed", gjson.GetBytes(event, "type").String())
|
||||
|
||||
require.Eventually(t, func() bool {
|
||||
logs := harness.moderationRepo.logSnapshot()
|
||||
return len(logs) == 1 && logs[0].Action == service.ContentModerationActionCyberPolicy &&
|
||||
strings.Contains(logs[0].Error, "upstream_usage=in:11,out:3")
|
||||
}, 3*time.Second, 10*time.Millisecond, "handler AfterTurn must call recordCyberPolicyIfMarked and write the risk-control event")
|
||||
|
||||
keyCtx, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
keyCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(requestPayload))
|
||||
blockKey := service.CyberSessionExplicitBlockKey(harness.apiKey.ID, keyCtx, []byte(requestPayload))
|
||||
require.NotEmpty(t, blockKey)
|
||||
store, ok := harness.gatewayCache.(service.CyberSessionBlockStore)
|
||||
require.True(t, ok)
|
||||
require.Eventually(t, func() bool {
|
||||
matched, findErr := store.FindCyberSessionBlocked(context.Background(), []string{blockKey})
|
||||
return findErr == nil && matched == blockKey
|
||||
}, 3*time.Second, 10*time.Millisecond, "handler AfterTurn must write the cyber session block table")
|
||||
|
||||
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"cyber-session-1","input":"follow-up"}`))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, _, err = harness.clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
var closeErr coderws.CloseError
|
||||
require.ErrorAs(t, err, &closeErr)
|
||||
require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code)
|
||||
// closeOpenAIClientWS caps close reasons at 120 bytes; passthrough must expose
|
||||
// the same client-visible prefix rather than dropping the close frame.
|
||||
require.Equal(t, "该会话已被网络安全策略屏蔽,请开启新会话 / This session is blocked by cyber-security policy, please ", closeErr.Reason)
|
||||
select {
|
||||
case <-harness.handlerDone:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("websocket handler did not exit")
|
||||
}
|
||||
select {
|
||||
case <-upstreamDone:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("upstream websocket did not exit")
|
||||
}
|
||||
select {
|
||||
case second := <-secondUpstreamFrame:
|
||||
t.Fatalf("blocked follow-up reached upstream: %s", second)
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesWebSocketV2PassthroughNonCyberTurnAllowsFollowup(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
upstreamDone := make(chan struct{})
|
||||
secondUpstreamFrame := make(chan []byte, 1)
|
||||
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
defer close(upstreamDone)
|
||||
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = conn.CloseNow() }()
|
||||
|
||||
readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
_, _, err = conn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, err)
|
||||
|
||||
firstCompleted := []byte(`{"type":"response.completed","response":{"id":"resp_non_cyber_handler_turn_1","model":"gpt-5.1","usage":{"input_tokens":2,"output_tokens":1}}}`)
|
||||
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
err = conn.Write(writeCtx, coderws.MessageText, firstCompleted)
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second)
|
||||
_, second, err := conn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, err)
|
||||
secondUpstreamFrame <- append([]byte(nil), second...)
|
||||
|
||||
secondCompleted := []byte(`{"type":"response.completed","response":{"id":"resp_non_cyber_handler_turn_2","model":"gpt-5.1","usage":{"input_tokens":3,"output_tokens":1}}}`)
|
||||
writeCtx, cancelWrite = context.WithTimeout(r.Context(), 3*time.Second)
|
||||
err = conn.Write(writeCtx, coderws.MessageText, secondCompleted)
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead = context.WithTimeout(r.Context(), 3*time.Second)
|
||||
_, _, _ = conn.Read(readCtx)
|
||||
cancelRead()
|
||||
}))
|
||||
defer upstreamServer.Close()
|
||||
harness := newOpenAIWSPassthroughHandlerHarness(t, upstreamServer.URL)
|
||||
|
||||
firstPayload := `{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"non-cyber-session-1","input":"first"}`
|
||||
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err := harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(firstPayload))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, firstEvent, err := harness.clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "resp_non_cyber_handler_turn_1", gjson.GetBytes(firstEvent, "response.id").String())
|
||||
|
||||
secondPayload := `{"type":"response.create","model":"gpt-5.1","prompt_cache_key":"non-cyber-session-1","input":"follow-up"}`
|
||||
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = harness.clientConn.Write(writeCtx, coderws.MessageText, []byte(secondPayload))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, secondEvent, err := harness.clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "resp_non_cyber_handler_turn_2", gjson.GetBytes(secondEvent, "response.id").String())
|
||||
require.Empty(t, harness.moderationRepo.logSnapshot())
|
||||
|
||||
keyCtx, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||
keyCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(firstPayload))
|
||||
blockKey := service.CyberSessionExplicitBlockKey(harness.apiKey.ID, keyCtx, []byte(firstPayload))
|
||||
require.NotEmpty(t, blockKey)
|
||||
store, ok := harness.gatewayCache.(service.CyberSessionBlockStore)
|
||||
require.True(t, ok)
|
||||
matched, findErr := store.FindCyberSessionBlocked(context.Background(), []string{blockKey})
|
||||
require.NoError(t, findErr)
|
||||
require.Empty(t, matched)
|
||||
|
||||
require.NoError(t, harness.clientConn.Close(coderws.StatusNormalClosure, "done"))
|
||||
select {
|
||||
case <-harness.handlerDone:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("non-cyber websocket handler did not exit")
|
||||
}
|
||||
select {
|
||||
case <-upstreamDone:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("non-cyber upstream websocket did not exit")
|
||||
}
|
||||
select {
|
||||
case second := <-secondUpstreamFrame:
|
||||
require.JSONEq(t, secondPayload, string(second))
|
||||
default:
|
||||
t.Fatal("non-cyber follow-up did not reach upstream")
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -19,13 +20,16 @@ type responsesFailedError struct {
|
||||
|
||||
// responsesFailedBody 对齐 apicompat.makeResponsesCompletedEvent 输出的 response 子对象字段集。
|
||||
// Output 用空 slice(不是 nil)确保 marshal 为 `[]` 而非 `null`。
|
||||
// CreatedAt 不带 omitempty:严格客户端把它当必填字段,缺失会以
|
||||
// `missing field 'created_at'` 反序列化失败——那正是本文件要避免的"客户端读不懂终止事件"。
|
||||
type responsesFailedBody struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
Model string `json:"model,omitempty"`
|
||||
Status string `json:"status"`
|
||||
Output []any `json:"output"`
|
||||
Error responsesFailedError `json:"error"`
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
Model string `json:"model,omitempty"`
|
||||
Status string `json:"status"`
|
||||
Output []any `json:"output"`
|
||||
Error responsesFailedError `json:"error"`
|
||||
}
|
||||
|
||||
// responsesFailedEvent 是写入 SSE data 行的顶层结构。
|
||||
@@ -61,11 +65,12 @@ func writeResponsesFailedSSE(c *gin.Context, errType, message string) bool {
|
||||
payload, err := json.Marshal(responsesFailedEvent{
|
||||
Type: "response.failed",
|
||||
Response: responsesFailedBody{
|
||||
ID: synthesizeResponseID(c),
|
||||
Object: "response",
|
||||
Model: requestModel(c),
|
||||
Status: "failed",
|
||||
Output: []any{},
|
||||
ID: synthesizeResponseID(c),
|
||||
Object: "response",
|
||||
CreatedAt: time.Now().Unix(),
|
||||
Model: requestModel(c),
|
||||
Status: "failed",
|
||||
Output: []any{},
|
||||
Error: responsesFailedError{
|
||||
Code: mapResponsesErrorCode(errType),
|
||||
Message: message,
|
||||
|
||||
@@ -227,6 +227,22 @@ func TestOpenAIHandleStreamingAwareError_BareResponsesRouteEmitsResponseFailed(t
|
||||
}
|
||||
|
||||
// Synthesized response.failed id falls back to uuid when no request_id is present.
|
||||
// issue #5601:严格的 Responses 客户端把 created_at 当必填字段,缺失即
|
||||
// `missing field 'created_at'`。合成的终止事件若解析不了,本文件存在的意义
|
||||
// (给客户端一个可识别的终止事件而不是盲重连)就落空了。
|
||||
func TestOpenAIHandleStreamingAwareError_ResponsesStreamingCarriesCreatedAt(t *testing.T) {
|
||||
c, w := newGinContextForEndpoint(t, EndpointResponses)
|
||||
h := &OpenAIGatewayHandler{}
|
||||
h.handleStreamingAwareError(c, http.StatusBadGateway, "upstream_error", "boom", true)
|
||||
|
||||
resp, _ := parseResponsesFailedSSE(t, w.Body.String())
|
||||
raw, ok := resp["created_at"]
|
||||
assert.True(t, ok, "response.failed 必须带 created_at")
|
||||
createdAt, ok := raw.(float64)
|
||||
assert.True(t, ok, "created_at 必须是数字,得到 %T", raw)
|
||||
assert.Greater(t, int64(createdAt), int64(0), "created_at 必须是有效的 unix 时间戳")
|
||||
}
|
||||
|
||||
func TestSynthesizeResponseID_FallbackUUID(t *testing.T) {
|
||||
c, _ := newGinContextForEndpoint(t, EndpointResponses)
|
||||
id := synthesizeResponseID(c)
|
||||
|
||||
@@ -21,10 +21,13 @@ func AnthropicToResponsesResponse(resp *AnthropicResponse) *ResponsesResponse {
|
||||
id = generateResponsesID()
|
||||
}
|
||||
|
||||
// Anthropic responses carry no creation timestamp, so stamp now — the same
|
||||
// synthesize-what-the-client-requires rule the generated id above follows.
|
||||
out := &ResponsesResponse{
|
||||
ID: id,
|
||||
Object: "response",
|
||||
Model: resp.Model,
|
||||
ID: id,
|
||||
Object: "response",
|
||||
CreatedAt: time.Now().Unix(),
|
||||
Model: resp.Model,
|
||||
}
|
||||
|
||||
var outputs []ResponsesOutput
|
||||
@@ -551,11 +554,12 @@ func makeResponsesCreatedEvent(state *AnthropicEventToResponsesState) ResponsesS
|
||||
Type: "response.created",
|
||||
SequenceNumber: seq,
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
Model: state.Model,
|
||||
Status: "in_progress",
|
||||
Output: []ResponsesOutput{},
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: "in_progress",
|
||||
Output: []ResponsesOutput{},
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -602,6 +606,7 @@ func makeResponsesCompletedEvent(
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: status,
|
||||
Output: outputs,
|
||||
|
||||
@@ -249,8 +249,8 @@ func anthropicUserToChatMessages(raw json.RawMessage) ([]ChatMessage, error) {
|
||||
|
||||
// anthropicAssistantToChatMessages handles an Anthropic assistant message.
|
||||
// Text content → assistant message content; tool_use blocks → tool_calls on the
|
||||
// same assistant message; thinking blocks are dropped (Chat Completions has no
|
||||
// inbound thinking field, matching anthropicAssistantToResponses).
|
||||
// same assistant message; thinking blocks → reasoning_content, but only on a
|
||||
// message that carries tool calls (see anthropicThinkingToReasoningContent).
|
||||
func anthropicAssistantToChatMessages(raw json.RawMessage) ([]ChatMessage, error) {
|
||||
// Plain string → single assistant message.
|
||||
var s string
|
||||
@@ -289,9 +289,40 @@ func anthropicAssistantToChatMessages(raw json.RawMessage) ([]ChatMessage, error
|
||||
})
|
||||
}
|
||||
|
||||
msg.ReasoningContent = anthropicThinkingToReasoningContent(blocks, len(msg.ToolCalls) > 0)
|
||||
|
||||
return []ChatMessage{msg}, nil
|
||||
}
|
||||
|
||||
// anthropicThinkingToReasoningContent folds thinking blocks back into the
|
||||
// Chat Completions reasoning_content field.
|
||||
//
|
||||
// chatMessageToAnthropicBlocks emits the upstream's reasoning_content as a
|
||||
// thinking block on the way out, so a multi-turn client echoes it back on the
|
||||
// next request; dropping it here made the bridge lose exactly what it had just
|
||||
// produced. DeepSeek's thinking mode requires the reasoning_content that
|
||||
// produced a tool call to be replayed on that assistant message and answers
|
||||
// 400 otherwise, which is why buildChatMessagesFromItems already carries
|
||||
// pendingReasoning onto assistant tool-call messages in the Responses→Chat
|
||||
// bridge. hasToolCalls keeps the scope identical to that sibling: reasoning
|
||||
// rides along with tool calls only, never on a plain assistant text turn.
|
||||
//
|
||||
// redacted_thinking blocks and signature-only placeholders carry no plaintext
|
||||
// and contribute nothing. Multiple blocks join with "\n", matching
|
||||
// extractResponsesReasoningText.
|
||||
func anthropicThinkingToReasoningContent(blocks []AnthropicContentBlock, hasToolCalls bool) string {
|
||||
if !hasToolCalls {
|
||||
return ""
|
||||
}
|
||||
var parts []string
|
||||
for _, b := range blocks {
|
||||
if b.Type == "thinking" && b.Thinking != "" {
|
||||
parts = append(parts, b.Thinking)
|
||||
}
|
||||
}
|
||||
return strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
// anthropicToolsToChatTools maps Anthropic tool definitions to Chat Completions
|
||||
// function tools. Server-side tools (web_search_*) are dropped — they have no
|
||||
// Chat Completions equivalent.
|
||||
|
||||
@@ -143,8 +143,11 @@ func TestAnthropicToChatCompletionsRequest_ThinkingDropped(t *testing.T) {
|
||||
out, err := AnthropicToChatCompletionsRequest(req)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Messages, 1)
|
||||
// Only text survives; thinking is dropped
|
||||
// Only text survives. Thinking is dropped because this turn carries no tool
|
||||
// calls — reasoning rides along with tool calls only, matching the
|
||||
// Responses→Chat bridge (see anthropicThinkingToReasoningContent).
|
||||
require.Equal(t, `"answer"`, string(out.Messages[0].Content))
|
||||
require.Empty(t, out.Messages[0].ReasoningContent)
|
||||
}
|
||||
|
||||
func TestAnthropicToChatCompletionsRequest_ToolChoiceAuto(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,229 @@
|
||||
package apicompat
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// issue #5528:/v1/messages 客户端(Claude Code 等)打到只会 Chat Completions 的
|
||||
// OpenAI 兼容上游时,历史 assistant 消息里的 thinking 块被整块丢弃。DeepSeek 的
|
||||
// thinking mode 要求产生工具调用的 reasoning_content 随该 assistant 消息回传,
|
||||
// 于是「单轮正常、一进多轮工具对话必现 400」。
|
||||
|
||||
func anthropicAssistantMsg(t *testing.T, blocks string) *AnthropicRequest {
|
||||
t.Helper()
|
||||
return &AnthropicRequest{
|
||||
Model: "deepseek-v4-flash",
|
||||
MaxTokens: 256,
|
||||
Messages: []AnthropicMessage{
|
||||
{Role: "user", Content: json.RawMessage(`"what's the weather?"`)},
|
||||
{Role: "assistant", Content: json.RawMessage(blocks)},
|
||||
{Role: "user", Content: json.RawMessage(`[{"type":"tool_result","tool_use_id":"toolu_1","content":"sunny"}]`)},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
const anthropicThinkingToolTurn = `[
|
||||
{"type":"thinking","thinking":"user wants weather, call the tool"},
|
||||
{"type":"text","text":"checking"},
|
||||
{"type":"tool_use","id":"toolu_1","name":"get_weather","input":{"city":"SF"}}
|
||||
]`
|
||||
|
||||
func TestAnthropicToChatCompletionsRequest_ThinkingBecomesReasoningContentOnToolTurn(t *testing.T) {
|
||||
out, err := AnthropicToChatCompletionsRequest(anthropicAssistantMsg(t, anthropicThinkingToolTurn))
|
||||
require.NoError(t, err)
|
||||
|
||||
var assistant *ChatMessage
|
||||
for i := range out.Messages {
|
||||
if out.Messages[i].Role == "assistant" {
|
||||
assistant = &out.Messages[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, assistant, "assistant message must survive the bridge")
|
||||
require.Equal(t, "user wants weather, call the tool", assistant.ReasoningContent,
|
||||
"产生工具调用的 thinking 必须作为 reasoning_content 回传,否则 DeepSeek 400")
|
||||
require.Len(t, assistant.ToolCalls, 1)
|
||||
require.Equal(t, `"checking"`, string(assistant.Content), "text/tool_use 处理保持不变")
|
||||
}
|
||||
|
||||
// 上游线格式才是上游看到的东西:字段没序列化出去,等于没修。
|
||||
func TestAnthropicToChatCompletionsRequest_ReasoningContentSerializesOnWire(t *testing.T) {
|
||||
out, err := AnthropicToChatCompletionsRequest(anthropicAssistantMsg(t, anthropicThinkingToolTurn))
|
||||
require.NoError(t, err)
|
||||
|
||||
payload, err := json.Marshal(out)
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, string(payload), `"reasoning_content":"user wants weather, call the tool"`)
|
||||
}
|
||||
|
||||
// 闭环不变式:thinking 块本来就是本桥出站时用上游 reasoning_content 生成的
|
||||
// (chatMessageToAnthropicBlocks),客户端只是原样回传。出站造、入站丢 = 自己丢自己的东西。
|
||||
func TestAnthropicChatBridge_ReasoningSurvivesOutboundInboundRoundTrip(t *testing.T) {
|
||||
upstream := ChatMessage{
|
||||
Role: "assistant",
|
||||
ReasoningContent: "step 1: need the weather tool",
|
||||
Content: json.RawMessage(`"checking"`),
|
||||
ToolCalls: []ChatToolCall{{
|
||||
ID: "call_1",
|
||||
Type: "function",
|
||||
Function: ChatFunctionCall{Name: "get_weather", Arguments: `{"city":"SF"}`},
|
||||
}},
|
||||
}
|
||||
|
||||
// 出站:Chat 响应 → Anthropic content blocks
|
||||
blocks := chatMessageToAnthropicBlocks(upstream)
|
||||
require.Equal(t, "thinking", blocks[0].Type)
|
||||
require.Equal(t, upstream.ReasoningContent, blocks[0].Thinking)
|
||||
|
||||
// 客户端下一轮把同一组 blocks 原样回传
|
||||
raw, err := json.Marshal(blocks)
|
||||
require.NoError(t, err)
|
||||
|
||||
// 入站:Anthropic content blocks → Chat 请求
|
||||
back, err := anthropicAssistantToChatMessages(raw)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, back, 1)
|
||||
require.Equal(t, upstream.ReasoningContent, back[0].ReasoningContent,
|
||||
"出站生成的 thinking 必须能原样还原回 reasoning_content")
|
||||
require.Len(t, back[0].ToolCalls, 1)
|
||||
}
|
||||
|
||||
// 兄弟不变式:Responses→Chat 桥(buildChatMessagesFromItems 的 pendingReasoning)
|
||||
// 早就把 reasoning 挂到带 tool_calls 的 assistant 消息上了。等价历史下两条桥必须一致。
|
||||
func TestAnthropicChatBridge_MatchesResponsesChatBridgeReasoningPlacement(t *testing.T) {
|
||||
responsesReq := &ResponsesRequest{
|
||||
Model: "deepseek-v4-flash",
|
||||
Input: json.RawMessage(`[
|
||||
{"type":"message","role":"user","content":[{"type":"input_text","text":"what's the weather?"}]},
|
||||
{"type":"reasoning","summary":[{"type":"summary_text","text":"call the tool"}]},
|
||||
{"type":"function_call","call_id":"call_1","name":"get_weather","arguments":"{\"city\":\"SF\"}"},
|
||||
{"type":"function_call_output","call_id":"call_1","output":"sunny"}
|
||||
]`),
|
||||
}
|
||||
viaResponses, err := ResponsesToChatCompletionsRequest(responsesReq)
|
||||
require.NoError(t, err)
|
||||
|
||||
viaAnthropic, err := AnthropicToChatCompletionsRequest(anthropicAssistantMsg(t, `[
|
||||
{"type":"thinking","thinking":"call the tool"},
|
||||
{"type":"tool_use","id":"toolu_1","name":"get_weather","input":{"city":"SF"}}
|
||||
]`))
|
||||
require.NoError(t, err)
|
||||
|
||||
reasoningOnToolCallMessage := func(msgs []ChatMessage) string {
|
||||
for _, m := range msgs {
|
||||
if m.Role == "assistant" && len(m.ToolCalls) > 0 {
|
||||
return m.ReasoningContent
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
require.Equal(t, "call the tool", reasoningOnToolCallMessage(viaResponses.Messages),
|
||||
"前置条件:兄弟桥本来就带 reasoning_content")
|
||||
require.Equal(t, reasoningOnToolCallMessage(viaResponses.Messages),
|
||||
reasoningOnToolCallMessage(viaAnthropic.Messages),
|
||||
"两条桥对等价历史必须产出同样的 reasoning_content 位置")
|
||||
}
|
||||
|
||||
// 作用域守卫:不带工具调用的纯文本轮次维持现状(与兄弟桥一致 —— reasoning 只随
|
||||
// 工具调用回传),避免把 reasoning_content 撒到不需要它的上游请求上。
|
||||
func TestAnthropicToChatCompletionsRequest_ThinkingWithoutToolCallsStaysDropped(t *testing.T) {
|
||||
req := &AnthropicRequest{
|
||||
Model: "deepseek-v4-flash",
|
||||
MaxTokens: 100,
|
||||
Messages: []AnthropicMessage{
|
||||
{Role: "assistant", Content: json.RawMessage(
|
||||
`[{"type":"thinking","thinking":"secret thoughts"},{"type":"text","text":"answer"}]`)},
|
||||
},
|
||||
}
|
||||
|
||||
out, err := AnthropicToChatCompletionsRequest(req)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, out.Messages, 1)
|
||||
require.Empty(t, out.Messages[0].ReasoningContent)
|
||||
require.Equal(t, `"answer"`, string(out.Messages[0].Content))
|
||||
|
||||
payload, err := json.Marshal(out)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, string(payload), "reasoning_content")
|
||||
}
|
||||
|
||||
func TestAnthropicThinkingToReasoningContent(t *testing.T) {
|
||||
blocksOf := func(t *testing.T, raw string) []AnthropicContentBlock {
|
||||
t.Helper()
|
||||
var blocks []AnthropicContentBlock
|
||||
require.NoError(t, json.Unmarshal([]byte(raw), &blocks))
|
||||
return blocks
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
raw string
|
||||
hasToolCalls bool
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "single_thinking_block",
|
||||
raw: `[{"type":"thinking","thinking":"a"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "a",
|
||||
},
|
||||
{
|
||||
// 多个 thinking 块用 "\n" 连接,与 extractResponsesReasoningText 一致。
|
||||
name: "multiple_blocks_join_with_newline",
|
||||
raw: `[{"type":"thinking","thinking":"a"},{"type":"text","text":"x"},{"type":"thinking","thinking":"b"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "a\nb",
|
||||
},
|
||||
{
|
||||
// redacted_thinking 没有明文可回传。
|
||||
name: "redacted_thinking_has_no_plaintext",
|
||||
raw: `[{"type":"redacted_thinking","signature":"abc"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
// 只带 signature 的 thinking 占位块(xAI/Codex 密文回放形态)同样无明文。
|
||||
name: "signature_only_thinking",
|
||||
raw: `[{"type":"thinking","thinking":"","signature":"gAAAAxxx"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "no_tool_calls_returns_empty",
|
||||
raw: `[{"type":"thinking","thinking":"a"}]`,
|
||||
hasToolCalls: false,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "no_thinking_blocks",
|
||||
raw: `[{"type":"text","text":"x"}]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "empty_blocks",
|
||||
raw: `[]`,
|
||||
hasToolCalls: true,
|
||||
want: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
require.Equal(t, tc.want,
|
||||
anthropicThinkingToReasoningContent(blocksOf(t, tc.raw), tc.hasToolCalls))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 纯字符串形态的 assistant content 没有 blocks 可读,走早返回分支,不得 panic。
|
||||
func TestAnthropicAssistantToChatMessages_PlainStringContentUnaffected(t *testing.T) {
|
||||
msgs, err := anthropicAssistantToChatMessages(json.RawMessage(`"just text"`))
|
||||
require.NoError(t, err)
|
||||
require.Len(t, msgs, 1)
|
||||
require.Empty(t, msgs[0].ReasoningContent)
|
||||
require.Equal(t, `"just text"`, string(msgs[0].Content))
|
||||
}
|
||||
@@ -1211,9 +1211,20 @@ func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model str
|
||||
id = generateResponsesID()
|
||||
}
|
||||
|
||||
// Carry the upstream's own creation timestamp when it sent one; otherwise
|
||||
// stamp now, same fallback shape as the generated id above.
|
||||
createdAt := int64(0)
|
||||
if resp != nil {
|
||||
createdAt = resp.Created
|
||||
}
|
||||
if createdAt <= 0 {
|
||||
createdAt = time.Now().Unix()
|
||||
}
|
||||
|
||||
out := &ResponsesResponse{
|
||||
ID: id,
|
||||
Object: "response",
|
||||
CreatedAt: createdAt,
|
||||
Model: model,
|
||||
Status: "completed",
|
||||
ServiceTier: chatServiceTier(resp),
|
||||
@@ -1711,6 +1722,7 @@ func FinalizeChatCompletionsResponsesStream(state *ChatCompletionsToResponsesStr
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: status,
|
||||
ServiceTier: state.ServiceTier,
|
||||
@@ -1731,6 +1743,7 @@ func ensureChatToResponsesCreated(state *ChatCompletionsToResponsesStreamState)
|
||||
Response: &ResponsesResponse{
|
||||
ID: state.ResponseID,
|
||||
Object: "response",
|
||||
CreatedAt: state.Created,
|
||||
Model: state.Model,
|
||||
Status: "in_progress",
|
||||
ServiceTier: state.ServiceTier,
|
||||
|
||||
@@ -358,6 +358,9 @@ func normalizeClientToolOutput(item map[string]any) {
|
||||
if _, ok := output.(string); ok {
|
||||
return
|
||||
}
|
||||
if isResponsesToolOutputContent(output) {
|
||||
return
|
||||
}
|
||||
if output == nil {
|
||||
item["output"] = ""
|
||||
return
|
||||
@@ -370,6 +373,25 @@ func normalizeClientToolOutput(item map[string]any) {
|
||||
item["output"] = string(encoded)
|
||||
}
|
||||
|
||||
func isResponsesToolOutputContent(output any) bool {
|
||||
parts, ok := output.([]any)
|
||||
if !ok || len(parts) == 0 {
|
||||
return false
|
||||
}
|
||||
for _, part := range parts {
|
||||
typed, ok := part.(map[string]any)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
switch stringValue(typed["type"]) {
|
||||
case "input_text", "input_image", "input_file":
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// normalizeToolSearchOutput converts both tool_search output wire shapes into
|
||||
// the string output required by function_call_output. Older clients send an
|
||||
// output field directly; newer Codex clients return discovered definitions in
|
||||
|
||||
@@ -460,7 +460,55 @@ func TestAdaptResponsesClientToolsWithInheritedMapping_LowersFollowupHistoryWith
|
||||
output := requireResponsesClientToolValue[map[string]any](t, items[1])
|
||||
require.Equal(t, "function_call_output", output["type"])
|
||||
require.NotContains(t, output, "id")
|
||||
require.JSONEq(t, `[{"text":"ok","type":"input_text"}]`, requireResponsesClientToolValue[string](t, output["output"]))
|
||||
require.Equal(t, []any{map[string]any{"type": "input_text", "text": "ok"}}, output["output"])
|
||||
}
|
||||
|
||||
func TestAdaptResponsesClientTools_NormalizesCustomToolOutput(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
output any
|
||||
wantOutput any
|
||||
}{
|
||||
{
|
||||
name: "supported content parts remain an array",
|
||||
output: []any{
|
||||
map[string]any{"type": "input_text", "text": "ok"},
|
||||
map[string]any{"type": "input_image", "image_url": "https://example.com/image.png"},
|
||||
map[string]any{"type": "input_file", "file_id": "file_123"},
|
||||
},
|
||||
wantOutput: []any{
|
||||
map[string]any{"type": "input_text", "text": "ok"},
|
||||
map[string]any{"type": "input_image", "image_url": "https://example.com/image.png"},
|
||||
map[string]any{"type": "input_file", "file_id": "file_123"},
|
||||
},
|
||||
},
|
||||
{name: "ordinary object is stringified", output: map[string]any{"ok": true}, wantOutput: `{"ok":true}`},
|
||||
{name: "arbitrary array is stringified", output: []any{"ok"}, wantOutput: `["ok"]`},
|
||||
{name: "empty array is stringified", output: []any{}, wantOutput: `[]`},
|
||||
{name: "mixed array is stringified", output: []any{map[string]any{"type": "input_text", "text": "ok"}, "bad"}, wantOutput: `[{"text":"ok","type":"input_text"},"bad"]`},
|
||||
{name: "unknown content type is stringified", output: []any{map[string]any{"type": "output_text", "text": "bad"}}, wantOutput: `[{"text":"bad","type":"output_text"}]`},
|
||||
{name: "whitespace-padded content type is stringified", output: []any{map[string]any{"type": " input_text ", "text": "bad"}}, wantOutput: `[{"text":"bad","type":" input_text "}]`},
|
||||
{name: "missing content type is stringified", output: []any{map[string]any{"text": "bad"}}, wantOutput: `[{"text":"bad"}]`},
|
||||
{name: "non-string content type is stringified", output: []any{map[string]any{"type": 1}}, wantOutput: `[{"type":1}]`},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
req := map[string]any{
|
||||
"tools": []any{map[string]any{"type": "custom", "name": "exec"}},
|
||||
"input": []any{map[string]any{
|
||||
"type": "custom_tool_call_output", "call_id": "call_1", "output": tc.output,
|
||||
}},
|
||||
}
|
||||
|
||||
_, changed, err := AdaptResponsesClientTools(req)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
item := requireResponsesClientToolValue[map[string]any](t, requireResponsesClientToolValue[[]any](t, req["input"])[0])
|
||||
require.Equal(t, "function_call_output", item["type"])
|
||||
require.Equal(t, tc.wantOutput, item["output"])
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAdaptResponsesClientToolsWithInheritedMapping_PromotesOmittedToolsDiscoveryIntoEffectiveDeclarations(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
package apicompat
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// issue #5601:严格的 Responses 客户端(Rust serde 系,如 Codex / Grok CLI)把
|
||||
// created_at 声明为必填字段,缺失即 `missing field 'created_at'` 反序列化失败。
|
||||
// 网关合成的 Responses 对象(Chat→Responses、Anthropic→Responses 两座桥)此前从不
|
||||
// 写这个字段——尽管两个流式 state 早就采集好了 Created 时间戳,只是没有出口。
|
||||
// 原生 Responses 透传走 gjson/sjson 字节级改写,不受影响。
|
||||
|
||||
// responseObjectOf 取出事件里的 response 子对象(按线格式,而不是按 Go 结构体)。
|
||||
func responseObjectOf(t *testing.T, evt ResponsesStreamEvent) map[string]any {
|
||||
t.Helper()
|
||||
m := marshalEvent(t, evt)
|
||||
resp, ok := m["response"].(map[string]any)
|
||||
require.True(t, ok, "event must carry a response object: %v", m)
|
||||
return resp
|
||||
}
|
||||
|
||||
func requireCreatedAt(t *testing.T, resp map[string]any) int64 {
|
||||
t.Helper()
|
||||
raw, ok := resp["created_at"]
|
||||
require.True(t, ok, "response 对象必须带 created_at,否则严格客户端直接反序列化失败")
|
||||
value, ok := raw.(float64)
|
||||
require.True(t, ok, "created_at 必须是数字,得到 %T", raw)
|
||||
require.Greater(t, int64(value), int64(0), "created_at 必须是有效的 unix 时间戳")
|
||||
return int64(value)
|
||||
}
|
||||
|
||||
// omitempty 陷阱守卫:created_at 为 0 时也必须出现在线格式里,
|
||||
// 否则「字段存在」这件事就依赖于运行时恰好非零。
|
||||
func TestWire_CreatedAtPresentEvenAtZero(t *testing.T) {
|
||||
resp := responseObjectOf(t, ResponsesStreamEvent{
|
||||
Type: "response.created",
|
||||
Response: &ResponsesResponse{ID: "resp_1", Object: "response", Status: "in_progress"},
|
||||
})
|
||||
require.Contains(t, resp, "created_at", "created_at 不得带 omitempty")
|
||||
require.EqualValues(t, 0, resp["created_at"])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Chat Completions → Responses
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestChatCompletionsResponseToResponses_CarriesCreatedAt(t *testing.T) {
|
||||
t.Run("uses_upstream_created_when_present", func(t *testing.T) {
|
||||
out := ChatCompletionsResponseToResponses(&ChatCompletionsResponse{
|
||||
ID: "chatcmpl_1",
|
||||
Created: 1700000000,
|
||||
Model: "deepseek-v4-flash",
|
||||
Choices: []ChatChoice{{Message: ChatMessage{Role: "assistant", Content: json.RawMessage(`"hi"`)}}},
|
||||
}, "deepseek-v4-flash", nil, nil, false, nil)
|
||||
require.EqualValues(t, 1700000000, out.CreatedAt, "上游给了 created 就照搬,不要另起时间")
|
||||
})
|
||||
|
||||
t.Run("stamps_now_when_upstream_omits_created", func(t *testing.T) {
|
||||
out := ChatCompletionsResponseToResponses(&ChatCompletionsResponse{
|
||||
ID: "chatcmpl_2",
|
||||
Model: "deepseek-v4-flash",
|
||||
Choices: []ChatChoice{{Message: ChatMessage{Role: "assistant", Content: json.RawMessage(`"hi"`)}}},
|
||||
}, "deepseek-v4-flash", nil, nil, false, nil)
|
||||
require.Greater(t, out.CreatedAt, int64(0))
|
||||
})
|
||||
|
||||
t.Run("nil_upstream_response_still_stamps", func(t *testing.T) {
|
||||
out := ChatCompletionsResponseToResponses(nil, "deepseek-v4-flash", nil, nil, false, nil)
|
||||
require.Greater(t, out.CreatedAt, int64(0), "空上游响应也必须产出可解析的对象")
|
||||
})
|
||||
}
|
||||
|
||||
// 同一条流里 response.created 与终止事件必须报同一个 created_at
|
||||
// (官方语义:created_at 是这次 response 的创建时刻,不随事件变化)。
|
||||
func TestChatCompletionsToResponsesStream_CreatedAtStableAcrossEvents(t *testing.T) {
|
||||
state := NewChatCompletionsToResponsesStreamState("deepseek-v4-flash")
|
||||
require.Greater(t, state.Created, int64(0), "前提:state 早就采集了时间戳")
|
||||
|
||||
var chunk ChatCompletionsChunk
|
||||
require.NoError(t, json.Unmarshal(
|
||||
[]byte(`{"choices":[{"index":0,"delta":{"content":"hi"}}]}`), &chunk))
|
||||
|
||||
events := ChatCompletionsChunkToResponsesEvents(&chunk, state)
|
||||
events = append(events, FinalizeChatCompletionsResponsesStream(state)...)
|
||||
|
||||
seen := map[string]int64{}
|
||||
for _, evt := range events {
|
||||
if evt.Response == nil {
|
||||
continue
|
||||
}
|
||||
seen[evt.Type] = requireCreatedAt(t, responseObjectOf(t, evt))
|
||||
}
|
||||
|
||||
require.Contains(t, seen, "response.created")
|
||||
require.Contains(t, seen, "response.completed")
|
||||
require.Equal(t, state.Created, seen["response.created"])
|
||||
require.Equal(t, seen["response.created"], seen["response.completed"],
|
||||
"同一条流的 created_at 必须恒定")
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Anthropic → Responses
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestAnthropicToResponsesResponse_StampsCreatedAt(t *testing.T) {
|
||||
out := AnthropicToResponsesResponse(&AnthropicResponse{
|
||||
ID: "msg_1",
|
||||
Type: "message",
|
||||
Role: "assistant",
|
||||
Model: "claude-sonnet-4-20250514",
|
||||
Content: []AnthropicContentBlock{{Type: "text", Text: "hi"}},
|
||||
})
|
||||
require.Greater(t, out.CreatedAt, int64(0),
|
||||
"Anthropic 响应不带时间戳,网关必须自己盖一个")
|
||||
}
|
||||
|
||||
func TestAnthropicEventToResponsesStream_CreatedAtStableAcrossEvents(t *testing.T) {
|
||||
state := NewAnthropicEventToResponsesState()
|
||||
state.Model = "claude-sonnet-4-20250514"
|
||||
require.Greater(t, state.Created, int64(0), "前提:state 早就采集了时间戳")
|
||||
|
||||
var events []ResponsesStreamEvent
|
||||
for _, raw := range []string{
|
||||
`{"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-20250514","content":[],"usage":{"input_tokens":3,"output_tokens":0}}}`,
|
||||
`{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`,
|
||||
`{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"hi"}}`,
|
||||
`{"type":"content_block_stop","index":0}`,
|
||||
`{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":2}}`,
|
||||
`{"type":"message_stop"}`,
|
||||
} {
|
||||
var evt AnthropicStreamEvent
|
||||
require.NoError(t, json.Unmarshal([]byte(raw), &evt))
|
||||
events = append(events, AnthropicEventToResponsesEvents(&evt, state)...)
|
||||
}
|
||||
events = append(events, FinalizeAnthropicResponsesStream(state)...)
|
||||
|
||||
seen := map[string]int64{}
|
||||
for _, evt := range events {
|
||||
if evt.Response == nil {
|
||||
continue
|
||||
}
|
||||
seen[evt.Type] = requireCreatedAt(t, responseObjectOf(t, evt))
|
||||
}
|
||||
|
||||
require.Contains(t, seen, "response.created")
|
||||
require.Contains(t, seen, "response.completed")
|
||||
require.Equal(t, state.Created, seen["response.created"])
|
||||
require.Equal(t, seen["response.created"], seen["response.completed"],
|
||||
"同一条流的 created_at 必须恒定")
|
||||
}
|
||||
|
||||
// ResponsesClientToolStreamRestorer 对部分事件走 unmarshal→re-marshal。
|
||||
// 结构体没有该字段时,上游带来的 created_at 会在这一步被静默抹掉。
|
||||
func TestResponsesStreamEvent_CreatedAtSurvivesUnmarshalRemarshal(t *testing.T) {
|
||||
upstream := []byte(`{"type":"response.completed","response":{"id":"resp_9","object":"response",` +
|
||||
`"created_at":1700000123,"model":"gpt-5.5","status":"completed","output":[]}}`)
|
||||
|
||||
var evt ResponsesStreamEvent
|
||||
require.NoError(t, json.Unmarshal(upstream, &evt))
|
||||
require.EqualValues(t, 1700000123, evt.Response.CreatedAt)
|
||||
|
||||
require.EqualValues(t, 1700000123, requireCreatedAt(t, responseObjectOf(t, evt)))
|
||||
}
|
||||
@@ -358,8 +358,13 @@ func (t *ResponsesTool) UnmarshalJSON(data []byte) error {
|
||||
|
||||
// ResponsesResponse is the non-streaming response from POST /v1/responses.
|
||||
type ResponsesResponse struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"` // "response"
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"` // "response"
|
||||
// CreatedAt is the unix creation timestamp. Strict Responses clients declare
|
||||
// it non-optional and abort with `missing field 'created_at'` when it is
|
||||
// absent, so it is always emitted — no omitempty. Same rule as ID (see the
|
||||
// "clients treat it as required" fallback in ChatCompletionsResponseToAnthropic).
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
Model string `json:"model"`
|
||||
Status string `json:"status"` // "completed" | "incomplete" | "failed"
|
||||
Output []ResponsesOutput `json:"output"`
|
||||
|
||||
@@ -133,6 +133,17 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont
|
||||
return false
|
||||
}
|
||||
|
||||
// Self-built images requests always carry a matching image_generation tool, so a
|
||||
// "tool choice not found in 'tools'" 400 means upstream revoked this account's
|
||||
// image capability. Gated on the self-built marker: passthrough clients control
|
||||
// their own tools/tool_choice and could otherwise poison a healthy account.
|
||||
if isOpenAIImagesSelfBuiltRequest(ctx) && isOpenAIImageCapabilityLossError(statusCode, responseBody) {
|
||||
if s != nil && s.rateLimitService != nil {
|
||||
_ = s.rateLimitService.HandleOpenAIImageCapabilityLoss(stateCtx, account, statusCode, responseBody)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
if s == nil || account == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -477,7 +477,7 @@ func configuredCodexGrokReasoningLevels(modelID string) []configuredCodexReasoni
|
||||
{Effort: "medium", Description: "Balanced reasoning for most coding tasks"},
|
||||
{Effort: "high", Description: "Greater reasoning depth for coding and agent tasks"},
|
||||
}
|
||||
if grokSupportsXHighReasoningEffort(modelID) {
|
||||
if GrokSupportsXHighReasoningEffort(modelID) {
|
||||
levels = append(levels, configuredCodexReasoningLevel{
|
||||
Effort: "xhigh",
|
||||
Description: "Extra-high reasoning depth for difficult tasks",
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
@@ -117,8 +118,11 @@ func writeOpenAICompactSSEFailureMessage(c *gin.Context, statusCode int, errType
|
||||
"response": map[string]any{
|
||||
"id": "resp_" + strings.ReplaceAll(uuid.NewString(), "-", ""),
|
||||
"object": "response",
|
||||
"status": "failed",
|
||||
"output": []any{},
|
||||
// 严格客户端把 created_at 当必填字段,缺失会反序列化失败,
|
||||
// 终止事件就白发了(退化成盲重连)。与 writeResponsesFailedSSE 对齐。
|
||||
"created_at": time.Now().Unix(),
|
||||
"status": "failed",
|
||||
"output": []any{},
|
||||
"error": map[string]any{
|
||||
"code": errType,
|
||||
"message": message,
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
//go:build unit
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// issue #5601:严格的 Responses 客户端把 created_at 当必填字段,缺失即
|
||||
// `missing field 'created_at'`。writeOpenAICompactSSEFailureMessage 存在的理由就是
|
||||
// 让 Codex 能把这帧识别成合法终止事件;解析不了就退化回它想避免的盲重连。
|
||||
func TestWriteOpenAICompactSSEFailureMessage_CarriesCreatedAt(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
|
||||
|
||||
writeOpenAICompactSSEFailureMessage(c, http.StatusBadGateway, "upstream_error", "boom")
|
||||
|
||||
body := rec.Body.String()
|
||||
require.Contains(t, body, "event: response.failed")
|
||||
|
||||
_, payload, found := strings.Cut(body, "data: ")
|
||||
require.True(t, found, "SSE 帧必须带 data 行: %q", body)
|
||||
|
||||
var event struct {
|
||||
Type string `json:"type"`
|
||||
Response struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
Status string `json:"status"`
|
||||
} `json:"response"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal([]byte(strings.TrimSpace(payload)), &event))
|
||||
|
||||
require.Equal(t, "response.failed", event.Type)
|
||||
require.Equal(t, "response", event.Response.Object)
|
||||
require.Equal(t, "failed", event.Response.Status)
|
||||
require.Greater(t, event.Response.CreatedAt, int64(0),
|
||||
"response.failed 必须带有效的 created_at,否则严格客户端读不出这帧")
|
||||
}
|
||||
@@ -285,6 +285,7 @@ func (s *OpenAIGatewayService) streamRawChatCompletions(
|
||||
clientOutputStarted := false
|
||||
pendingLines := make([]string, 0, 8)
|
||||
refusalDetector := newOpenAIChatSilentRefusalDetector(requestBodyLen)
|
||||
var terminal openAIRawStreamTerminalState
|
||||
|
||||
writeLine := func(line string) {
|
||||
if clientDisconnected {
|
||||
@@ -323,6 +324,7 @@ func (s *OpenAIGatewayService) streamRawChatCompletions(
|
||||
refusalDetector.ObserveSSELine(line)
|
||||
if payload, ok := extractOpenAISSEDataLine(line); ok {
|
||||
trimmedPayload := strings.TrimSpace(payload)
|
||||
terminal.ObserveDataLine(trimmedPayload)
|
||||
if trimmedPayload != "[DONE]" {
|
||||
observer.ObserveOpenAI([]byte(payload), strings.TrimSpace(gjson.Get(payload, "type").String()))
|
||||
usageOnlyChunk := isOpenAIChatUsageOnlyStreamChunk(payload)
|
||||
@@ -350,14 +352,65 @@ func (s *OpenAIGatewayService) streamRawChatCompletions(
|
||||
}
|
||||
}
|
||||
|
||||
if err := scanner.Err(); err != nil {
|
||||
if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
|
||||
logger.L().Warn("openai chat_completions raw: stream read error",
|
||||
zap.Error(err),
|
||||
zap.String("request_id", requestID),
|
||||
)
|
||||
resultWithUsage := func() *OpenAIForwardResult {
|
||||
return &OpenAIForwardResult{
|
||||
RequestID: requestID,
|
||||
Usage: usage,
|
||||
Model: originalModel,
|
||||
BillingModel: billingModel,
|
||||
UpstreamModel: upstreamModel,
|
||||
UpstreamResponseModel: observedUpstreamResponseModel(c),
|
||||
UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c),
|
||||
UpstreamResponseServiceTier: observedUpstreamResponseServiceTier(c),
|
||||
ReasoningEffort: reasoningEffort,
|
||||
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
|
||||
Stream: true,
|
||||
Duration: time.Since(startTime),
|
||||
FirstTokenMs: firstTokenMs,
|
||||
}
|
||||
} else if !clientDisconnected && !clientOutputStarted {
|
||||
}
|
||||
|
||||
scanErr := scanner.Err()
|
||||
if scanErr != nil && !errors.Is(scanErr, context.Canceled) && !errors.Is(scanErr, context.DeadlineExceeded) {
|
||||
logger.L().Warn("openai chat_completions raw: stream read error",
|
||||
zap.Error(scanErr),
|
||||
zap.String("request_id", requestID),
|
||||
)
|
||||
}
|
||||
|
||||
// 客户端取消/断开后上游读失败与上游截断不可区分(取消会连带取消上游请求),
|
||||
// 沿用既有语义:按已收到的用量正常收尾计费,不判为上游故障。
|
||||
clientAborted := clientDisconnected ||
|
||||
errors.Is(scanErr, context.Canceled) ||
|
||||
errors.Is(scanErr, context.DeadlineExceeded)
|
||||
|
||||
// 上游在任何终止信号之前结束:连接被 reset(scanErr != nil)或干净 EOF。
|
||||
// 两者都不能再记成功——此前统一返回 nil error,把上游截断伪装成
|
||||
// `HTTP 200 + usage 0/0`,客户端收到半截回答且 Ops 侧完全无感。
|
||||
if !clientAborted && terminal.IsTruncated(clientOutputStarted) {
|
||||
cause := scanErr
|
||||
if cause == nil {
|
||||
cause = ErrOpenAIUpstreamStreamTruncated
|
||||
}
|
||||
logger.L().Warn("openai chat_completions raw: upstream stream truncated before terminal chunk",
|
||||
zap.Error(cause),
|
||||
zap.String("request_id", requestID),
|
||||
zap.Int64("account_id", account.ID),
|
||||
zap.String("upstream_model", upstreamModel),
|
||||
zap.Bool("saw_sse_data", terminal.sawDataLine),
|
||||
zap.Bool("client_output_started", clientOutputStarted),
|
||||
)
|
||||
if !clientOutputStarted {
|
||||
// 响应头尚未提交:可以透明换号重试,客户端不会看到半截流。
|
||||
return nil, newOpenAIRawStreamTruncatedFailoverError(c, account, requestID, cause)
|
||||
}
|
||||
// 已写出语义字节:无法再 failover,改为带类型的上游错误。handler 会据此
|
||||
// 补发 SSE error 帧并把本次请求计入 SLA 失败。
|
||||
recordOpenAIRawStreamTruncation(c, account, requestID, cause, "http_error")
|
||||
return resultWithUsage(), newOpenAIUpstreamStreamReadError(cause)
|
||||
}
|
||||
|
||||
if scanErr == nil && !clientDisconnected && !clientOutputStarted {
|
||||
if refusalDetector.IsSilentRefusal() {
|
||||
return nil, newOpenAISilentRefusalFailoverError(c, account, requestID)
|
||||
}
|
||||
@@ -380,21 +433,7 @@ func (s *OpenAIGatewayService) streamRawChatCompletions(
|
||||
}
|
||||
}
|
||||
|
||||
return &OpenAIForwardResult{
|
||||
RequestID: requestID,
|
||||
Usage: usage,
|
||||
Model: originalModel,
|
||||
BillingModel: billingModel,
|
||||
UpstreamModel: upstreamModel,
|
||||
UpstreamResponseModel: observedUpstreamResponseModel(c),
|
||||
UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c),
|
||||
UpstreamResponseServiceTier: observedUpstreamResponseServiceTier(c),
|
||||
ReasoningEffort: reasoningEffort,
|
||||
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
|
||||
Stream: true,
|
||||
Duration: time.Since(startTime),
|
||||
FirstTokenMs: firstTokenMs,
|
||||
}, nil
|
||||
return resultWithUsage(), nil
|
||||
}
|
||||
|
||||
// ensureOpenAIChatStreamUsage 确保 raw Chat Completions 流式请求会让上游返回 usage。
|
||||
|
||||
@@ -687,6 +687,326 @@ func TestForwardAsRawChatCompletions_StripsEmptyToolCallIdentity(t *testing.T) {
|
||||
require.True(t, followUpSeen)
|
||||
}
|
||||
|
||||
// 上游在生成中途干净 EOF(无 [DONE]/usage/finish_reason)且已向客户端写出内容:
|
||||
// 不能再记成 HTTP 200 成功,必须回带类型化的上游截断错误,由 handler 补 SSE error
|
||||
// 帧并计入 SLA 失败。
|
||||
func TestForwardAsRawChatCompletions_TruncatedStreamAfterOutputFailsRequest(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-pro","messages":[{"role":"user","content":"hello"}],"stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
upstreamBody := strings.Join([]string{
|
||||
`data: {"id":"chatcmpl_cut","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[{"index":0,"delta":{"content":"half an ans"},"finish_reason":null}]}`,
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_truncated"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||
}}
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "")
|
||||
require.Error(t, err)
|
||||
require.NotNil(t, result, "已收字节的用量仍需带回,供 ops 记录首 token 时延")
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.False(t, errors.As(err, &failoverErr), "已写出语义字节后不得再 failover")
|
||||
|
||||
code, message, ok := OpenAIUpstreamStreamReadErrorDetails(err)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, OpenAIUpstreamStreamTruncatedCode, code)
|
||||
require.NotEmpty(t, message)
|
||||
// 已写出的内容保持原样透传,客户端拿到的仍是它已经收到的那部分。
|
||||
require.Contains(t, rec.Body.String(), `"content":"half an ans"`)
|
||||
require.NotContains(t, rec.Body.String(), "data: [DONE]")
|
||||
}
|
||||
|
||||
// 上游 200 但一个 SSE 字节都没发:响应头尚未提交,应换号重试而不是回 200 空流。
|
||||
func TestForwardAsRawChatCompletions_EmptyStreamBeforeOutputTriggersFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-pro","messages":[{"role":"user","content":"hello"}],"stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", 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{"text/event-stream"}, "x-request-id": []string{"rid_empty"}},
|
||||
Body: io.NopCloser(strings.NewReader("")),
|
||||
}}
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "")
|
||||
require.Nil(t, result)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.True(t, errors.As(err, &failoverErr))
|
||||
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
|
||||
require.Equal(t, OpenAIUpstreamStreamTruncatedCode,
|
||||
gjson.GetBytes(failoverErr.ResponseBody, "error.code").String())
|
||||
require.True(t, failoverErr.ShouldRetryNextAccount())
|
||||
require.False(t, c.Writer.Written(), "换号重试前不得提交 200 响应头")
|
||||
require.Empty(t, rec.Body.String())
|
||||
}
|
||||
|
||||
// 传输层错误(Cloudflare edge reset 等)在写出后同样不能记成功,且分类要区别于
|
||||
// 干净 EOF,便于 ops 分辨 reset 与静默截断。
|
||||
func TestForwardAsRawChatCompletions_StreamReadErrorAfterOutputFailsRequest(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-pro","messages":[{"role":"user","content":"hello"}],"stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", 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{"text/event-stream"}, "x-request-id": []string{"rid_reset"}},
|
||||
Body: &openAIChatStreamReadErrorCloser{
|
||||
payload: []byte(`data: {"id":"chatcmpl_reset","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[{"index":0,"delta":{"content":"partial"}}]}` + "\n\n"),
|
||||
err: errors.New("read tcp 172.18.0.4->172.65.90.23:443: read: connection reset by peer"),
|
||||
},
|
||||
}}
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "")
|
||||
require.Error(t, err)
|
||||
require.NotNil(t, result)
|
||||
|
||||
code, _, ok := OpenAIUpstreamStreamReadErrorDetails(err)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, OpenAIUpstreamStreamReadErrorCode, code)
|
||||
require.Contains(t, rec.Body.String(), `"content":"partial"`)
|
||||
}
|
||||
|
||||
// 边界:缺 [DONE] 但收到了 usage 帧 —— 生成已完整,只是尾巴丢失。必须继续按成功
|
||||
// 计费,否则会误伤那些跑完就直接 EOF 的兼容上游并白送 token。
|
||||
func TestForwardAsRawChatCompletions_MissingDoneWithUsageStillSucceeds(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-pro","messages":[{"role":"user","content":"hello"}],"stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
upstreamBody := strings.Join([]string{
|
||||
`data: {"id":"chatcmpl_nodone","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[{"index":0,"delta":{"content":"ok"}}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_nodone","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[],"usage":{"prompt_tokens":11,"completion_tokens":6,"total_tokens":17}}`,
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_nodone_usage"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||
}}
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, 11, result.Usage.InputTokens)
|
||||
require.Equal(t, 6, result.Usage.OutputTokens)
|
||||
}
|
||||
|
||||
// 边界:缺 [DONE] 与 usage,但末帧带 finish_reason —— 生成正常结束,同样不判截断。
|
||||
func TestForwardAsRawChatCompletions_MissingDoneWithFinishReasonStillSucceeds(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-pro","messages":[{"role":"user","content":"hello"}],"stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
upstreamBody := strings.Join([]string{
|
||||
`data: {"id":"chatcmpl_finish","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[{"index":0,"delta":{"content":"done"},"finish_reason":null}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_finish","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}`,
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_nodone_finish"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||
}}
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Contains(t, rec.Body.String(), `"finish_reason":"stop"`)
|
||||
}
|
||||
|
||||
// openAIRawStreamDisconnectedWriter 模拟客户端已断开:raw 直转路径经
|
||||
// WriteString 写出,故两个方法都必须失败(只覆盖 Write 会被内嵌 writer 绕过)。
|
||||
type openAIRawStreamDisconnectedWriter struct {
|
||||
gin.ResponseWriter
|
||||
}
|
||||
|
||||
func (w *openAIRawStreamDisconnectedWriter) Write([]byte) (int, error) {
|
||||
return 0, errors.New("write failed: client disconnected")
|
||||
}
|
||||
|
||||
func (w *openAIRawStreamDisconnectedWriter) WriteString(string) (int, error) {
|
||||
return 0, errors.New("write failed: client disconnected")
|
||||
}
|
||||
|
||||
// 客户端已断开时上游随后截断:两者不可区分,沿用既有语义按已收用量正常收尾计费,
|
||||
// 不得把客户端离场记成上游故障。
|
||||
func TestForwardAsRawChatCompletions_ClientDisconnectTruncationStillBills(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-pro","messages":[{"role":"user","content":"hello"}],"stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Writer = &openAIRawStreamDisconnectedWriter{ResponseWriter: c.Writer}
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
upstreamBody := strings.Join([]string{
|
||||
`data: {"id":"chatcmpl_gone","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[{"index":0,"delta":{"content":"ok"}}]}`,
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_gone"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||
}}
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
}
|
||||
|
||||
// 客户端取消会连带取消上游请求,上游读因此报 context.Canceled:同样不判为上游截断。
|
||||
func TestForwardAsRawChatCompletions_ClientCancelTruncationStillBills(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-pro","messages":[{"role":"user","content":"hello"}],"stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", 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{"text/event-stream"}, "x-request-id": []string{"rid_cancel"}},
|
||||
Body: &openAIChatStreamReadErrorCloser{
|
||||
payload: []byte(`data: {"id":"chatcmpl_cancel","object":"chat.completion.chunk","model":"deepseek-v4-pro","choices":[{"index":0,"delta":{"content":"ok"}}]}` + "\n\n"),
|
||||
err: context.Canceled,
|
||||
},
|
||||
}}
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.forwardAsRawChatCompletions(context.Background(), c, rawChatCompletionsTestAccount(), body, "")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
}
|
||||
|
||||
func TestOpenAIRawStreamTerminalState(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
payloads []string
|
||||
clientStarted bool
|
||||
wantTerminated bool
|
||||
wantTruncated bool
|
||||
}{
|
||||
{
|
||||
name: "done sentinel",
|
||||
payloads: []string{`{"choices":[{"delta":{"content":"a"}}]}`, "[DONE]"},
|
||||
clientStarted: true,
|
||||
wantTerminated: true,
|
||||
},
|
||||
{
|
||||
name: "usage chunk",
|
||||
payloads: []string{`{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1}}`},
|
||||
clientStarted: true,
|
||||
wantTerminated: true,
|
||||
},
|
||||
{
|
||||
name: "finish reason",
|
||||
payloads: []string{`{"choices":[{"delta":{},"finish_reason":"length"}]}`},
|
||||
clientStarted: true,
|
||||
wantTerminated: true,
|
||||
},
|
||||
{
|
||||
name: "null finish reason is not terminal",
|
||||
payloads: []string{`{"choices":[{"delta":{"content":"a"},"finish_reason":null}]}`},
|
||||
clientStarted: true,
|
||||
wantTruncated: true,
|
||||
},
|
||||
{
|
||||
name: "usage null is not terminal",
|
||||
payloads: []string{`{"choices":[{"delta":{"content":"a"}}],"usage":null}`},
|
||||
clientStarted: true,
|
||||
wantTruncated: true,
|
||||
},
|
||||
{
|
||||
// 上游对 stream 请求回了裸 JSON:无 data: 行,既有行为是原样透传。
|
||||
name: "non-sse body already forwarded",
|
||||
payloads: nil,
|
||||
clientStarted: true,
|
||||
},
|
||||
{
|
||||
name: "no bytes at all",
|
||||
payloads: nil,
|
||||
clientStarted: false,
|
||||
wantTruncated: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
var state openAIRawStreamTerminalState
|
||||
for _, payload := range tt.payloads {
|
||||
state.ObserveDataLine(payload)
|
||||
}
|
||||
require.Equal(t, tt.wantTerminated, state.Terminated())
|
||||
require.Equal(t, tt.wantTruncated, state.IsTruncated(tt.clientStarted))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardAsRawChatCompletions_ClientDisconnectDrainsUsage(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
@@ -121,7 +121,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
}
|
||||
if shouldStripOpenAIResponsesInputNamespaces(account, wsDecision.Transport, passthroughEnabled) {
|
||||
keepToolCallNamespaces := shouldKeepOpenAIResponsesToolCallNamespaces(
|
||||
account, wsDecision.Transport, passthroughEnabled, compactPath,
|
||||
account, wsDecision.Transport, passthroughEnabled, compactPath, body,
|
||||
)
|
||||
body, err = stripOpenAIResponsesInputNamespaces(body, keepToolCallNamespaces)
|
||||
if err != nil {
|
||||
@@ -182,6 +182,17 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
return s.forwardResponsesViaRawChatCompletions(ctx, c, account, body)
|
||||
}
|
||||
if account.IsOpenAI() && (account.IsOpenAIApiKey() || account.IsOpenAIOAuthLike()) {
|
||||
normalizedReasoningBody, reasoningChanged, reasoningErr := normalizeOpenAIResponsesReasoningContentReplay(body)
|
||||
if reasoningErr != nil {
|
||||
return nil, fmt.Errorf("normalize OpenAI Responses reasoning content replay: %w", reasoningErr)
|
||||
}
|
||||
if reasoningChanged {
|
||||
body = normalizedReasoningBody
|
||||
originalBody = normalizedReasoningBody
|
||||
requestView = newOpenAIRequestView(normalizedReasoningBody)
|
||||
reqModel, reqStream, promptCacheKey = requestView.Model, requestView.Stream, requestView.PromptCacheKey
|
||||
originalModel = reqModel
|
||||
}
|
||||
sanitizedBody, changed, sanitizeErr := sanitizeOpenAIResponsesInputItemIDs(body)
|
||||
if sanitizeErr != nil {
|
||||
return nil, fmt.Errorf("sanitize OpenAI Responses input item IDs: %w", sanitizeErr)
|
||||
@@ -472,13 +483,29 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
if decodeErr != nil {
|
||||
return nil, decodeErr
|
||||
}
|
||||
// Responses OAuth 与 Chat 兼容入口保持一致:纯文本 system 可以无损提升后删除,
|
||||
// JSON object 模式仍需在 input 中保留 JSON 指令供上游兼容校验。
|
||||
omitPromotedSystemMessages := !strings.EqualFold(
|
||||
strings.TrimSpace(gjson.GetBytes(body, "text.format.type").String()),
|
||||
"json_object",
|
||||
)
|
||||
codexResult := codexTransformResult{}
|
||||
if compatMessagesBridge {
|
||||
codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{IsCodexCLI: isCodexCLI, IsCompact: isCompactRequest, SkipDefaultInstructions: true, PreserveToolCallIDs: true})
|
||||
codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{
|
||||
IsCodexCLI: isCodexCLI,
|
||||
IsCompact: isCompactRequest,
|
||||
SkipDefaultInstructions: true,
|
||||
PreserveToolCallIDs: true,
|
||||
OmitPromotedSystemMessagesFromInput: omitPromotedSystemMessages,
|
||||
})
|
||||
ensureCodexOAuthInstructionsField(decoded)
|
||||
markDecodedModified()
|
||||
} else {
|
||||
codexResult = applyCodexOAuthTransform(decoded, isCodexCLI, isCompactRequest)
|
||||
codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{
|
||||
IsCodexCLI: isCodexCLI,
|
||||
IsCompact: isCompactRequest,
|
||||
OmitPromotedSystemMessagesFromInput: omitPromotedSystemMessages,
|
||||
})
|
||||
}
|
||||
if codexResult.Error != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": codexResult.Error.Error()}})
|
||||
|
||||
@@ -723,7 +723,7 @@ func normalizeGrokReasoningEffortValue(raw, model string) (string, bool) {
|
||||
case "minimal":
|
||||
return "low", true
|
||||
case "xhigh", "extrahigh":
|
||||
if grokSupportsXHighReasoningEffort(model) {
|
||||
if GrokSupportsXHighReasoningEffort(model) {
|
||||
return "xhigh", true
|
||||
}
|
||||
return "high", true
|
||||
@@ -734,7 +734,9 @@ func normalizeGrokReasoningEffortValue(raw, model string) (string, bool) {
|
||||
}
|
||||
}
|
||||
|
||||
func grokSupportsXHighReasoningEffort(model string) bool {
|
||||
// GrokSupportsXHighReasoningEffort reports whether the model advertises and
|
||||
// forwards the xhigh reasoning effort (Grok 4.6 and its undated alias).
|
||||
func GrokSupportsXHighReasoningEffort(model string) bool {
|
||||
model = strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(model)))
|
||||
return model == "grok-4.6" || model == "grok-4.6-latest"
|
||||
}
|
||||
|
||||
@@ -458,6 +458,67 @@ func openAIRequestBodyHasTools(body []byte) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// normalizeOpenAIResponsesReasoningContentReplay removes non-portable
|
||||
// reasoning.content arrays before history is sent to a real OpenAI Responses
|
||||
// endpoint. Compatible providers may return visible reasoning blocks there,
|
||||
// while OpenAI accepts only an empty array when the item is replayed.
|
||||
//
|
||||
// Keep the reasoning item and its portable fields (summary, encrypted_content,
|
||||
// ids, and opaque extensions). Callers scope this normalization to OpenAI
|
||||
// destinations; compatible providers may still consume their own content.
|
||||
func normalizeOpenAIResponsesReasoningContentReplay(body []byte) ([]byte, bool, error) {
|
||||
input := gjson.GetBytes(body, "input")
|
||||
if !input.IsArray() {
|
||||
return body, false, nil
|
||||
}
|
||||
|
||||
needsNormalization := false
|
||||
input.ForEach(func(_, item gjson.Result) bool {
|
||||
if strings.TrimSpace(item.Get("type").String()) != "reasoning" {
|
||||
return true
|
||||
}
|
||||
content := item.Get("content")
|
||||
if content.IsArray() && len(content.Array()) > 0 {
|
||||
needsNormalization = true
|
||||
return false
|
||||
}
|
||||
return true
|
||||
})
|
||||
if !needsNormalization {
|
||||
return body, false, nil
|
||||
}
|
||||
|
||||
var reqBody map[string]any
|
||||
if err := decodeOpenAIJSONUseNumber(body, &reqBody); err != nil {
|
||||
return body, false, fmt.Errorf("normalize OpenAI reasoning content replay: %w", err)
|
||||
}
|
||||
items, ok := reqBody["input"].([]any)
|
||||
if !ok {
|
||||
return body, false, nil
|
||||
}
|
||||
changed := false
|
||||
for _, rawItem := range items {
|
||||
item, ok := rawItem.(map[string]any)
|
||||
if !ok || strings.TrimSpace(firstNonEmptyString(item["type"])) != "reasoning" {
|
||||
continue
|
||||
}
|
||||
content, ok := item["content"].([]any)
|
||||
if !ok || len(content) == 0 {
|
||||
continue
|
||||
}
|
||||
delete(item, "content")
|
||||
changed = true
|
||||
}
|
||||
if !changed {
|
||||
return body, false, nil
|
||||
}
|
||||
normalized, err := marshalOpenAIUpstreamJSON(reqBody)
|
||||
if err != nil {
|
||||
return body, false, fmt.Errorf("serialize normalized OpenAI reasoning content replay: %w", err)
|
||||
}
|
||||
return normalized, true, nil
|
||||
}
|
||||
|
||||
func normalizeOpenAIAPIKeyStoreFalseReasoningReplay(body []byte, knownStoreFalse bool) ([]byte, bool, error) {
|
||||
if !knownStoreFalse && gjson.GetBytes(body, "store").Type != gjson.False {
|
||||
return body, false, nil
|
||||
@@ -1018,6 +1079,12 @@ func normalizeOpenAIResponsesWebSocketCompatibilityBody(body []byte, account *Ac
|
||||
return body, false, err
|
||||
}
|
||||
}
|
||||
if next, normalizedReasoningContent, err := normalizeOpenAIResponsesReasoningContentReplay(normalized); err != nil {
|
||||
return body, false, err
|
||||
} else if normalizedReasoningContent {
|
||||
normalized = next
|
||||
changed = true
|
||||
}
|
||||
if account.IsOpenAIApiKey() {
|
||||
if next, normalizedParallel, err := normalizeOpenAIParallelToolCallsWithoutTools(normalized, responsesLite); err != nil {
|
||||
return body, false, err
|
||||
|
||||
@@ -353,3 +353,60 @@ func TestNormalizeOpenAIParallelToolCallsWithoutTools_KeepsResponsesLiteAddition
|
||||
require.False(t, changed)
|
||||
require.Equal(t, gjson.False, gjson.GetBytes(normalized, "parallel_tool_calls").Type)
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesReasoningContentReplayStripsCrossProviderArray(t *testing.T) {
|
||||
body := []byte(`{"model":"gpt-5.6-sol","input":[` +
|
||||
`{"type":"message","role":"user","content":"one"},` +
|
||||
`{"type":"message","role":"assistant","content":"two"},` +
|
||||
`{"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"},` +
|
||||
`{"type":"function_call_output","call_id":"call_1","output":"ok"},` +
|
||||
`{"type":"message","role":"user","content":"five"},` +
|
||||
`{"type":"reasoning","id":"rs_provider","summary":[{"type":"summary_text","text":"portable"}],"content":[{"type":"reasoning_text","text":"visible reasoning"}],"opaque":9007199254740993},` +
|
||||
`{"type":"message","role":"assistant","content":[{"type":"output_text","text":"answer"}]}` +
|
||||
`]}`)
|
||||
|
||||
normalized, changed, err := normalizeOpenAIResponsesReasoningContentReplay(body)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.Equal(t, "reasoning", gjson.GetBytes(normalized, "input.5.type").String())
|
||||
require.False(t, gjson.GetBytes(normalized, "input.5.content").Exists())
|
||||
require.Equal(t, "portable", gjson.GetBytes(normalized, "input.5.summary.0.text").String())
|
||||
require.Equal(t, "9007199254740993", gjson.GetBytes(normalized, "input.5.opaque").Raw)
|
||||
require.Equal(t, "answer", gjson.GetBytes(normalized, "input.6.content.0.text").String())
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesReasoningContentReplayKeepsPortableShapes(t *testing.T) {
|
||||
for _, body := range []string{
|
||||
`{"input":[{"type":"reasoning","summary":[]}]}`,
|
||||
`{"input":[{"type":"reasoning","content":[],"summary":[]}]}`,
|
||||
`{"input":[{"type":"message","content":[{"type":"input_text","text":"keep"}]}]}`,
|
||||
} {
|
||||
normalized, changed, err := normalizeOpenAIResponsesReasoningContentReplay([]byte(body))
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.JSONEq(t, body, string(normalized))
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyStripsReasoningContentOnlyForOpenAI(t *testing.T) {
|
||||
body := []byte(`{"type":"response.create","model":"gpt-5.6-sol","store":true,"input":[{"type":"reasoning","summary":[{"type":"summary_text","text":"keep"}],"content":[{"type":"reasoning_text","text":"remove"}]}]}`)
|
||||
for _, accountType := range []string{AccountTypeAPIKey, AccountTypeOAuth} {
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: accountType,
|
||||
}, false)
|
||||
require.NoError(t, err)
|
||||
require.True(t, changed)
|
||||
require.False(t, gjson.GetBytes(normalized, "input.0.content").Exists())
|
||||
require.Equal(t, "keep", gjson.GetBytes(normalized, "input.0.summary.0.text").String())
|
||||
}
|
||||
|
||||
normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{
|
||||
Platform: PlatformZhipu,
|
||||
Type: AccountTypeAPIKey,
|
||||
}, false)
|
||||
require.NoError(t, err)
|
||||
require.False(t, changed)
|
||||
require.JSONEq(t, string(body), string(normalized))
|
||||
}
|
||||
|
||||
@@ -250,6 +250,34 @@ func TestOpenAIPassthroughAPIKeyRestoresClientToolsNonStreaming(t *testing.T) {
|
||||
require.Equal(t, "*** Begin Patch", gjson.Get(recorder.Body.String(), "output.1.input").String())
|
||||
}
|
||||
|
||||
func TestOpenAIPassthroughAPIKeyPreservesCustomToolOutputContentParts(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{"model":"gpt-5.4","stream":false,"tools":[{"type":"custom","name":"exec"}],"input":[{"type":"custom_tool_call_output","call_id":"call_1","output":[{"type":"input_text","text":"result"},{"type":"input_file","file_id":"file_123"}]}]}`)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"resp_tools","status":"completed","output":[],"usage":{}}`)),
|
||||
}}
|
||||
svc := openAIClientToolsTestService(upstream)
|
||||
account := &Account{ID: 6240, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "test-key"}}
|
||||
|
||||
result, err := svc.forwardOpenAIPassthrough(context.Background(), c, account, body, body, "gpt-5.4", false, nil, false, time.Now())
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, "function_call_output", gjson.GetBytes(upstream.lastBody, "input.0.type").String())
|
||||
output := gjson.GetBytes(upstream.lastBody, "input.0.output")
|
||||
require.True(t, output.IsArray(), "native Responses content parts must reach the upstream as an array")
|
||||
require.Equal(t, "input_text", output.Get("0.type").String())
|
||||
require.Equal(t, "result", output.Get("0.text").String())
|
||||
require.Equal(t, "input_file", output.Get("1.type").String())
|
||||
require.Equal(t, "file_123", output.Get("1.file_id").String())
|
||||
}
|
||||
|
||||
func TestOpenAIPassthroughAPIKeyRestoresClientToolsStreaming(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := openAIClientToolsRequest(true)
|
||||
|
||||
@@ -37,6 +37,13 @@ type OpenAIImagesUpstreamError struct {
|
||||
Message string
|
||||
Param string
|
||||
UpstreamRequestID string
|
||||
|
||||
// SynthesizedFromModelText marks an error the gateway inferred from the
|
||||
// model's plain-text output instead of reading it off a structured upstream
|
||||
// error frame. Such a verdict describes this one turn ("the model answered
|
||||
// with words instead of an image"), not the account — see
|
||||
// shouldCoolOpenAIImagesToolForError.
|
||||
SynthesizedFromModelText bool
|
||||
}
|
||||
|
||||
func (e *OpenAIImagesUpstreamError) Error() string {
|
||||
@@ -328,6 +335,26 @@ func openAIImageUploadToDataURL(upload OpenAIImagesUpload) (string, error) {
|
||||
return "data:" + contentType + ";base64," + base64.StdEncoding.EncodeToString(upload.Data), nil
|
||||
}
|
||||
|
||||
// openAIImagesSelfBuiltRequestContextKey marks a request whose upstream body was
|
||||
// fully constructed by buildOpenAIImagesResponsesRequest, i.e. tool_choice and the
|
||||
// matching image_generation tool are always both present and never client-controlled.
|
||||
type openAIImagesSelfBuiltRequestContextKey struct{}
|
||||
|
||||
func withOpenAIImagesSelfBuiltRequest(ctx context.Context) context.Context {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
return context.WithValue(ctx, openAIImagesSelfBuiltRequestContextKey{}, true)
|
||||
}
|
||||
|
||||
func isOpenAIImagesSelfBuiltRequest(ctx context.Context) bool {
|
||||
if ctx == nil {
|
||||
return false
|
||||
}
|
||||
selfBuilt, _ := ctx.Value(openAIImagesSelfBuiltRequestContextKey{}).(bool)
|
||||
return selfBuilt
|
||||
}
|
||||
|
||||
func buildOpenAIImagesResponsesRequest(parsed *OpenAIImagesRequest, toolModel string) ([]byte, error) {
|
||||
if parsed == nil {
|
||||
return nil, fmt.Errorf("parsed images request is required")
|
||||
@@ -711,6 +738,10 @@ func openAIImagesTextFallbackErrorForText(text string) *OpenAIImagesUpstreamErro
|
||||
ErrorType: "upstream_error",
|
||||
Code: "image_generation_unavailable",
|
||||
Message: "Upstream did not execute image generation",
|
||||
// Inferred from the model's own words, not from an upstream error frame:
|
||||
// good enough to fail this turn over to another account, not evidence that
|
||||
// this account's image tool is down for the next 30 minutes.
|
||||
SynthesizedFromModelText: true,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1775,6 +1806,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth(
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
upstreamCtx = withOpenAIImagesSelfBuiltRequest(upstreamCtx)
|
||||
upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, true, parsed.StickySessionSeed(), false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1922,6 +1954,26 @@ const (
|
||||
openAIImagesOAuthUnavailableReason = "openai_images_oauth_tool_unavailable"
|
||||
)
|
||||
|
||||
// shouldCoolOpenAIImagesToolForError decides whether an image_generation_unavailable
|
||||
// verdict is durable enough to park the account's image tool for
|
||||
// openAIImagesOAuthUnavailableCooldown.
|
||||
//
|
||||
// Only an upstream error frame that names the condition qualifies. A verdict the
|
||||
// gateway synthesized from the model's plain-text reply does not: it merely says
|
||||
// this prompt produced words instead of an image, which is prompt-dependent and
|
||||
// happens on healthy accounts. Writing a 30-minute account-level cooldown from it
|
||||
// is doubly wrong because the very same error is classified retryable
|
||||
// (IsOpenAIImagesRetryableUpstreamError: status >= 500) and drives
|
||||
// newOpenAIAccountFailoverError — so one such reply walks the pool and cools every
|
||||
// account the retry touches.
|
||||
//
|
||||
// This mirrors the rule the alpha/search path already states in words: a
|
||||
// tool-endpoint failure "仍允许本次请求换号,但不修改任何账号状态"
|
||||
// (see shouldApplyOpenAIAlphaSearchAccountErrorSideEffects).
|
||||
func shouldCoolOpenAIImagesToolForError(upstreamErr *OpenAIImagesUpstreamError) bool {
|
||||
return upstreamErr != nil && !upstreamErr.SynthesizedFromModelText
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) coolOpenAIImagesOAuthTool(ctx context.Context, account *Account) {
|
||||
if s == nil || s.accountRepo == nil || account == nil || account.Platform != PlatformOpenAI {
|
||||
return
|
||||
@@ -2017,7 +2069,9 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthResponseError(
|
||||
|
||||
responseBody := openAIImagesUpstreamErrorResponseBody(upstreamErr)
|
||||
if upstreamErr.Code == "image_generation_unavailable" {
|
||||
s.coolOpenAIImagesOAuthTool(ctx, account)
|
||||
if shouldCoolOpenAIImagesToolForError(upstreamErr) {
|
||||
s.coolOpenAIImagesOAuthTool(ctx, account)
|
||||
}
|
||||
if responseWritten {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
//go:build unit
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// issue #6171:v0.1.181 起,/v1/images/generations 只要上游"回文字没回图",账号就被
|
||||
// 写 30 分钟 openai:image_generation 模型级冷却。该判据是**请求级**的(这个 prompt
|
||||
// 这一轮模型选择了说话),却被当成**账号级**能力失效;又因为同一个错误被判为
|
||||
// 可重试(502)并驱动 failover,一次闲聊回复会沿着号池逐个把账号冷却掉。
|
||||
|
||||
// countingModelRateLimitRepo 记录 SetModelRateLimit 调用,用于断言"没写账号状态"。
|
||||
type countingModelRateLimitRepo struct {
|
||||
accountRepoStub
|
||||
calls int
|
||||
scopes []string
|
||||
}
|
||||
|
||||
func (r *countingModelRateLimitRepo) SetModelRateLimit(_ context.Context, _ int64, scope string, _ time.Time, _ ...string) error {
|
||||
r.calls++
|
||||
r.scopes = append(r.scopes, scope)
|
||||
return nil
|
||||
}
|
||||
|
||||
func newImagesCooldownContext(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil)
|
||||
return c, rec
|
||||
}
|
||||
|
||||
func imagesCooldownAccount() *Account {
|
||||
return &Account{ID: 77, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "img-oauth"}
|
||||
}
|
||||
|
||||
func TestShouldCoolOpenAIImagesToolForError(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
err *OpenAIImagesUpstreamError
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "nil_error",
|
||||
err: nil,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
// 网关从模型文字里推断出来的判据:只说明这一轮没出图。
|
||||
name: "synthesized_from_model_text",
|
||||
err: &OpenAIImagesUpstreamError{
|
||||
StatusCode: http.StatusBadGateway,
|
||||
Code: "image_generation_unavailable",
|
||||
SynthesizedFromModelText: true,
|
||||
},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
// 上游自己在 error 帧里点名该状态:这才是账号级证据,保持冷却。
|
||||
name: "structured_upstream_error_frame",
|
||||
err: &OpenAIImagesUpstreamError{
|
||||
StatusCode: http.StatusBadGateway,
|
||||
Code: "image_generation_unavailable",
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
require.Equal(t, tc.want, shouldCoolOpenAIImagesToolForError(tc.err))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 主复现:文字兜底判据不得写账号级冷却。
|
||||
func TestHandleOpenAIImagesOAuthResponseError_TextFallbackDoesNotCoolAccount(t *testing.T) {
|
||||
c, _ := newImagesCooldownContext(t)
|
||||
repo := &countingModelRateLimitRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := imagesCooldownAccount()
|
||||
|
||||
upstreamErr := openAIImagesTextFallbackErrorForText("Here's a polished image prompt for your request.")
|
||||
require.NotNil(t, upstreamErr)
|
||||
require.Equal(t, "image_generation_unavailable", upstreamErr.Code)
|
||||
|
||||
err := svc.handleOpenAIImagesOAuthResponseError(
|
||||
context.Background(), c, account, "gpt-image-2", "https://upstream.example/v1/responses",
|
||||
&http.Response{StatusCode: http.StatusOK, Header: http.Header{}},
|
||||
OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c), upstreamErr,
|
||||
)
|
||||
|
||||
require.Zero(t, repo.calls, "模型闲聊不构成账号级证据,不得写 30 分钟冷却")
|
||||
|
||||
// 换号行为必须原样保留:本 PR 只撤销账号状态写入,不动 failover。
|
||||
var failover *UpstreamFailoverError
|
||||
require.True(t, errors.As(err, &failover), "仍应触发换号,got %T", err)
|
||||
}
|
||||
|
||||
// 对照不变式:上游 error 帧点名该状态时仍然冷却,否则等于把功能整个废掉。
|
||||
func TestHandleOpenAIImagesOAuthResponseError_StructuredUnavailableStillCoolsAccount(t *testing.T) {
|
||||
c, _ := newImagesCooldownContext(t)
|
||||
repo := &countingModelRateLimitRepo{}
|
||||
svc := &OpenAIGatewayService{accountRepo: repo}
|
||||
account := imagesCooldownAccount()
|
||||
|
||||
upstreamErr := &OpenAIImagesUpstreamError{
|
||||
StatusCode: http.StatusBadGateway,
|
||||
ErrorType: "upstream_error",
|
||||
Code: "image_generation_unavailable",
|
||||
Message: "image generation tool is not available for this account",
|
||||
}
|
||||
|
||||
_ = svc.handleOpenAIImagesOAuthResponseError(
|
||||
context.Background(), c, account, "gpt-image-2", "https://upstream.example/v1/responses",
|
||||
&http.Response{StatusCode: http.StatusOK, Header: http.Header{}},
|
||||
OpenAIImagesJSONKeepaliveAdjustedWrittenSize(c), upstreamErr,
|
||||
)
|
||||
|
||||
require.Equal(t, 1, repo.calls, "结构化上游证据仍须写冷却")
|
||||
require.Equal(t, []string{openAIImageGenerationRateLimitKey}, repo.scopes)
|
||||
}
|
||||
|
||||
// 标记必须打在文字兜底的两个入口上,且不影响违规拦截分支的判定。
|
||||
func TestOpenAIImagesTextFallback_MarksSynthesizedVerdicts(t *testing.T) {
|
||||
t.Run("plain_text_reply_is_synthesized", func(t *testing.T) {
|
||||
err := openAIImagesTextFallbackErrorForText("Here's a polished image prompt for your request.")
|
||||
require.NotNil(t, err)
|
||||
require.True(t, err.SynthesizedFromModelText)
|
||||
require.Equal(t, "image_generation_unavailable", err.Code)
|
||||
require.Equal(t, http.StatusBadGateway, err.StatusCode)
|
||||
})
|
||||
|
||||
t.Run("body_entrypoint_is_synthesized", func(t *testing.T) {
|
||||
body := []byte("event: response.completed\n" +
|
||||
`data: {"type":"response.completed","response":{"id":"r","status":"completed",` +
|
||||
`"output":[{"type":"message","content":[{"type":"output_text","text":"I drafted a prompt for you."}]}]}}` +
|
||||
"\n\n")
|
||||
err := openAIImagesTextFallbackError(body)
|
||||
require.NotNil(t, err)
|
||||
require.True(t, err.SynthesizedFromModelText)
|
||||
})
|
||||
|
||||
t.Run("content_policy_branch_unchanged", func(t *testing.T) {
|
||||
err := openAIImagesTextFallbackErrorForText("Blocked by our content policy.")
|
||||
require.NotNil(t, err)
|
||||
require.Equal(t, "content_policy_violation", err.Code)
|
||||
require.Equal(t, http.StatusBadRequest, err.StatusCode)
|
||||
// 该分支本来就不走冷却(Code 不匹配),标记与否都不改变行为;
|
||||
// 断言它没有被顺手打标,避免语义漂移。
|
||||
require.False(t, err.SynthesizedFromModelText)
|
||||
})
|
||||
|
||||
t.Run("empty_text_yields_no_error", func(t *testing.T) {
|
||||
require.Nil(t, openAIImagesTextFallbackErrorForText(" "))
|
||||
})
|
||||
}
|
||||
|
||||
// 级联的前提条件:该错误确实是可重试的,所以会带着"已写冷却"的副作用换号。
|
||||
// 这条用例把前提钉死,避免以后有人把 502 改成非重试后误以为本修复多余。
|
||||
func TestOpenAIImagesTextFallback_RemainsRetryableAndThusCascades(t *testing.T) {
|
||||
err := openAIImagesTextFallbackErrorForText("Here's a polished image prompt for your request.")
|
||||
require.NotNil(t, err)
|
||||
require.True(t, IsOpenAIImagesRetryableUpstreamError(err),
|
||||
"文字兜底判据是可重试的——正因如此,写账号冷却会沿号池级联")
|
||||
}
|
||||
@@ -132,6 +132,44 @@ func TestOpenAIGatewayService_ResponsesUnknownModelDoesNotFallbackToGPT54(t *tes
|
||||
require.True(t, rec.Code >= http.StatusBadRequest)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OAuthResponsesPromotesSystemMessageWithoutDuplication(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
const systemPrompt = "Unique system prefix for Responses token accounting."
|
||||
const existingInstructions = "Existing instructions."
|
||||
body := []byte(`{"model":"gpt-5.4","stream":false,"instructions":"` + existingInstructions + `","input":[{"role":"system","content":"` + systemPrompt + `"},{"role":"user","content":"hello"}]}`)
|
||||
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{err: errors.New("stop after capture")}
|
||||
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
|
||||
account := &Account{
|
||||
ID: 124,
|
||||
Name: "openai-oauth",
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "oauth-token",
|
||||
"chatgpt_account_id": "chatgpt-acc",
|
||||
},
|
||||
Status: StatusActive,
|
||||
Schedulable: true,
|
||||
}
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, account, body)
|
||||
|
||||
require.Error(t, err)
|
||||
require.Nil(t, result)
|
||||
require.NotEmpty(t, upstream.lastBody)
|
||||
require.Equal(t, systemPrompt+"\n\n"+existingInstructions, gjson.GetBytes(upstream.lastBody, "instructions").String())
|
||||
require.Equal(t, int64(1), gjson.GetBytes(upstream.lastBody, "input.#").Int())
|
||||
require.Equal(t, "user", gjson.GetBytes(upstream.lastBody, "input.0.role").String())
|
||||
require.Equal(t, 1, strings.Count(string(upstream.lastBody), systemPrompt))
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_NativeResponsesBodyModificationPreservesHTMLChars(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
// openAIRawStreamTruncatedUpstreamMessage 是 raw CC 直转路径上游截断的 ops 消息。
|
||||
const openAIRawStreamTruncatedUpstreamMessage = "Upstream Chat Completions stream ended before any terminal chunk"
|
||||
|
||||
// openAIRawStreamTerminalState 记录 raw Chat Completions SSE 流是否收到过
|
||||
// **终止信号**。
|
||||
//
|
||||
// 背景:CC 直转路径把上游 SSE 原样透传,此前只要 HTTP 状态是 200 就按成功收尾——
|
||||
// 上游中途断流(Cloudflare edge reset、后端 worker 掉线)会被伪装成
|
||||
// `HTTP 200 + usage 0/0`:客户端拿到半截回答,网关既不报错也不计入 SLA,
|
||||
// Ops 侧完全不可见。
|
||||
//
|
||||
// 三种终止信号任一出现即认为上游"讲完了",只是尾巴可能丢失,不作截断处理:
|
||||
//
|
||||
// - [DONE] —— OpenAI CC 协议标准哨兵
|
||||
// - usage chunk —— include_usage 生效时的末尾用量帧(网关强制打开)
|
||||
// - finish_reason —— 生成正常结束(stop/length/tool_calls/...)
|
||||
//
|
||||
// 只认 [DONE] 会误伤那些跑完最后一帧就直接 EOF 的兼容上游;只认 usage 会误伤
|
||||
// 不支持 include_usage 的上游。三者取并集,把误判压到"上游确实在生成中途被切断"。
|
||||
type openAIRawStreamTerminalState struct {
|
||||
// sawDataLine 表示上游至少发过一行 `data:`,即响应确实是 SSE 语义流。
|
||||
// 非 SSE 响应体(上游对 stream 请求返回裸 JSON)不参与截断判定,保持既有透传行为。
|
||||
sawDataLine bool
|
||||
sawDone bool
|
||||
sawUsage bool
|
||||
sawFinishReason bool
|
||||
}
|
||||
|
||||
// ObserveDataLine 从单行 SSE `data:` 载荷中提取终止信号。payload 需已 TrimSpace。
|
||||
func (t *openAIRawStreamTerminalState) ObserveDataLine(payload string) {
|
||||
if t == nil {
|
||||
return
|
||||
}
|
||||
t.sawDataLine = true
|
||||
if payload == "[DONE]" {
|
||||
t.sawDone = true
|
||||
return
|
||||
}
|
||||
if usage := gjson.Get(payload, "usage"); usage.Exists() && usage.IsObject() {
|
||||
t.sawUsage = true
|
||||
}
|
||||
if t.sawFinishReason {
|
||||
return
|
||||
}
|
||||
for _, choice := range gjson.Get(payload, "choices").Array() {
|
||||
// finish_reason 为 null 时 String() 返回空串,不算终止。
|
||||
if strings.TrimSpace(choice.Get("finish_reason").String()) != "" {
|
||||
t.sawFinishReason = true
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Terminated 表示上游给出过终止信号。
|
||||
func (t *openAIRawStreamTerminalState) Terminated() bool {
|
||||
return t != nil && (t.sawDone || t.sawUsage || t.sawFinishReason)
|
||||
}
|
||||
|
||||
// IsTruncated 判定上游是否在任何终止信号之前结束。clientOutputStarted 用于放行
|
||||
// 非 SSE 响应体:那类响应本就没有 data: 行,既有行为是原样透传,不在本次判定范围内;
|
||||
// 但"一个字节都没收到"的空 200 依然算截断。
|
||||
func (t *openAIRawStreamTerminalState) IsTruncated(clientOutputStarted bool) bool {
|
||||
if t == nil || t.Terminated() {
|
||||
return false
|
||||
}
|
||||
return t.sawDataLine || !clientOutputStarted
|
||||
}
|
||||
|
||||
// newOpenAIRawStreamTruncatedFailoverError 处理"上游截断且尚未向客户端写出任何
|
||||
// 字节"的情况:响应头还没提交,可以透明换号重试,客户端不会看到半截流。
|
||||
func newOpenAIRawStreamTruncatedFailoverError(
|
||||
c *gin.Context,
|
||||
account *Account,
|
||||
upstreamRequestID string,
|
||||
cause error,
|
||||
) *UpstreamFailoverError {
|
||||
recordOpenAIRawStreamTruncation(c, account, upstreamRequestID, cause, "failover")
|
||||
|
||||
headers := http.Header{}
|
||||
if id := strings.TrimSpace(upstreamRequestID); id != "" {
|
||||
headers.Set("x-request-id", id)
|
||||
}
|
||||
return &UpstreamFailoverError{
|
||||
StatusCode: http.StatusBadGateway,
|
||||
ResponseBody: openAIRawStreamTruncatedErrorBody(cause),
|
||||
ResponseHeaders: headers,
|
||||
}
|
||||
}
|
||||
|
||||
// recordOpenAIRawStreamTruncation 把上游截断记入 ops 上下文,使其在错误日志与
|
||||
// 账号健康度中可见——这正是此前"HTTP 200 假成功"丢掉的信息。
|
||||
func recordOpenAIRawStreamTruncation(
|
||||
c *gin.Context,
|
||||
account *Account,
|
||||
upstreamRequestID string,
|
||||
cause error,
|
||||
kind string,
|
||||
) {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
message := openAIRawStreamTruncatedMessage(cause)
|
||||
platform := PlatformOpenAI
|
||||
accountID := int64(0)
|
||||
accountName := ""
|
||||
if account != nil {
|
||||
platform = account.Platform
|
||||
accountID = account.ID
|
||||
accountName = account.Name
|
||||
}
|
||||
|
||||
setOpsUpstreamError(c, http.StatusBadGateway, message, "")
|
||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||
Platform: platform,
|
||||
AccountID: accountID,
|
||||
AccountName: accountName,
|
||||
UpstreamStatusCode: http.StatusBadGateway,
|
||||
UpstreamRequestID: strings.TrimSpace(upstreamRequestID),
|
||||
Kind: kind,
|
||||
Message: message,
|
||||
})
|
||||
}
|
||||
|
||||
// openAIRawStreamTruncatedMessage 拼出 ops 消息:干净 EOF 没有底层错误可带,
|
||||
// 传输层错误(connection reset / http2 stream error)则保留原因以便定位。
|
||||
func openAIRawStreamTruncatedMessage(cause error) string {
|
||||
if cause == nil || errors.Is(cause, ErrOpenAIUpstreamStreamTruncated) {
|
||||
return openAIRawStreamTruncatedUpstreamMessage
|
||||
}
|
||||
return openAIRawStreamTruncatedUpstreamMessage + ": " + cause.Error()
|
||||
}
|
||||
|
||||
// openAIRawStreamTruncatedErrorBody 构造 failover 错误体,code/message 与
|
||||
// 写出后走 openAIUpstreamStreamReadError 的客户端分类保持一致。
|
||||
func openAIRawStreamTruncatedErrorBody(cause error) []byte {
|
||||
code, message := classifyOpenAIUpstreamStreamReadError(cause)
|
||||
body, err := json.Marshal(map[string]any{
|
||||
"error": map[string]any{
|
||||
"type": "upstream_error",
|
||||
"code": code,
|
||||
"message": message,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return []byte(`{"error":{"type":"upstream_error","code":"` + OpenAIUpstreamStreamTruncatedCode +
|
||||
`","message":"Upstream response stream ended before completion"}}`)
|
||||
}
|
||||
return body
|
||||
}
|
||||
@@ -77,25 +77,48 @@ func shouldStripOpenAIResponsesInputNamespaces(account *Account, transport OpenA
|
||||
// 故 OAuth 非 compact 请求必须保留。
|
||||
// - compact 端点的 schema 不含该字段,携带即 400 `Unknown parameter:
|
||||
// input[N].namespace`(issue #4761 正文),故 compact 一律清理。
|
||||
// - API Key 出口是标准 Responses API(api.openai.com 或自定义 base_url),同样
|
||||
// 不认识该字段,维持全量清理;否则只能退化成
|
||||
// openai_responses_rejected_field_retry 的逐项删除,6 次上限根本盖不住长历史。
|
||||
// - API Key 出口默认按标准 Responses API 处理并清理该字段;但当请求本身声明
|
||||
// namespace 工具时,上游显然使用了 namespace 扩展,此时必须保留调用项上的
|
||||
// namespace,否则声明与历史调用会失配并触发 Missing namespace。
|
||||
// - 摊平模式下调用项已被改写成平名,残留 namespace 指向的声明已不存在,一律清理。
|
||||
func shouldKeepOpenAIResponsesToolCallNamespaces(
|
||||
account *Account,
|
||||
transport OpenAIUpstreamTransport,
|
||||
passthroughEnabled bool,
|
||||
compactPath bool,
|
||||
body []byte,
|
||||
) bool {
|
||||
if account == nil || !account.IsOpenAIOAuthLike() {
|
||||
if account == nil {
|
||||
return false
|
||||
}
|
||||
if compactPath {
|
||||
return false
|
||||
}
|
||||
if account.IsOpenAIApiKey() {
|
||||
return hasOpenAIResponsesNamespaceToolDeclaration(body)
|
||||
}
|
||||
if !account.IsOpenAIOAuthLike() {
|
||||
return false
|
||||
}
|
||||
return !shouldFlattenOpenAIResponsesNamespaces(account, transport, passthroughEnabled, compactPath)
|
||||
}
|
||||
|
||||
func hasOpenAIResponsesNamespaceToolDeclaration(body []byte) bool {
|
||||
tools := gjson.GetBytes(body, "tools")
|
||||
if !tools.IsArray() {
|
||||
return false
|
||||
}
|
||||
found := false
|
||||
tools.ForEach(func(_, tool gjson.Result) bool {
|
||||
if strings.EqualFold(strings.TrimSpace(tool.Get("type").String()), "namespace") {
|
||||
found = true
|
||||
return false
|
||||
}
|
||||
return true
|
||||
})
|
||||
return found
|
||||
}
|
||||
|
||||
// openAIResponsesToolCallItemTypes 是携带 namespace 的调用项类型集合。与
|
||||
// removeOpenAIResponsesRejectedNamespaceAtIndex 的反应式白名单保持一致;codex-rs
|
||||
// protocol/src/models.rs 中只有 FunctionCall 与 CustomToolCall 序列化 namespace,
|
||||
|
||||
@@ -66,6 +66,29 @@ func TestOpenAIGatewayService_OAuthPreservesCodexNamespaceTools(t *testing.T) {
|
||||
require.Empty(t, openAIResponsesNamespaceNames(c))
|
||||
}
|
||||
|
||||
// API Key 自定义上游若接受 namespace 工具声明,也要求历史 function_call 原样携带
|
||||
// namespace。声明仍为命名空间工具却清掉调用项字段,会触发 Missing namespace。
|
||||
func TestOpenAIGatewayService_APIKeyPreservesDeclaredNamespaceToolCalls(t *testing.T) {
|
||||
body := []byte(codexNamespaceRequestBody)
|
||||
upstream := &httpUpstreamRecorder{responses: []*http.Response{
|
||||
newOpenAIRejectedFieldTestResponse(http.StatusOK, namespaceForwardOKResponse),
|
||||
}}
|
||||
c := newOpenAIRejectedFieldTestContext(body)
|
||||
|
||||
result, err := newOpenAIRejectedFieldTestService(upstream).Forward(
|
||||
context.Background(), c, newOpenAIRejectedFieldTestAccount(), body,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Len(t, upstream.bodies, 1)
|
||||
forwarded := upstream.bodies[0]
|
||||
|
||||
require.True(t, gjson.GetBytes(forwarded, `tools.#(type=="namespace")`).Exists())
|
||||
require.Equal(t, "collaboration", gjson.GetBytes(forwarded, "input.0.namespace").String())
|
||||
require.False(t, gjson.GetBytes(forwarded, "input.1.namespace").Exists())
|
||||
}
|
||||
|
||||
// compact 端点 schema 更窄:input[].namespace 会 400 Unknown parameter(issue #4761),
|
||||
// 且没有证据表明它接受 namespace 工具声明。compact 只做历史摘要、不需要模型寻址工具,
|
||||
// 因此保持既有的摊平 + 全量清理行为,不随默认值翻转扩大风险面。
|
||||
|
||||
@@ -78,6 +78,7 @@ func TestShouldKeepOpenAIResponsesToolCallNamespaces(t *testing.T) {
|
||||
transport OpenAIUpstreamTransport
|
||||
passthroughEnabled bool
|
||||
compactPath bool
|
||||
body []byte
|
||||
want bool
|
||||
}{
|
||||
// 上游按 namespace 解析历史调用,缺字段会 400 "Missing namespace for function_call"。
|
||||
@@ -92,15 +93,20 @@ func TestShouldKeepOpenAIResponsesToolCallNamespaces(t *testing.T) {
|
||||
// WSv2 + compact 是唯一「不摊平但仍必须清理」的组合,钉住 compact 判定本身,
|
||||
// 使其不会被误当成可由 shouldFlatten 推导出的冗余分支。
|
||||
{name: "oauth_compact_wsv2_strips", account: oauth, transport: OpenAIUpstreamTransportResponsesWebsocketV2, compactPath: true, want: false},
|
||||
// API Key 出口是标准 Responses API,不认识该字段。
|
||||
{name: "apikey_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, want: false},
|
||||
// API Key 默认按标准 Responses API 清理;请求显式声明 namespace 工具时,
|
||||
// 自定义上游需要原样接收对应的历史调用。
|
||||
{name: "apikey_without_namespace_tool_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, want: false},
|
||||
{name: "apikey_with_namespace_tool_keeps", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, body: []byte(`{"tools":[{"type":"namespace","name":"mcp__codex_app","tools":[]}]}`), want: true},
|
||||
{name: "apikey_with_mixed_case_namespace_tool_keeps", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, body: []byte(`{"tools":[{"type":" Namespace ","name":"mcp__codex_app","tools":[]}]}`), want: true},
|
||||
{name: "apikey_function_tool_with_namespace_field_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, body: []byte(`{"tools":[{"type":"function","name":"automation_update","namespace":"mcp__codex_app"}]}`), want: false},
|
||||
{name: "apikey_compact_with_namespace_tool_strips", account: apiKey, transport: OpenAIUpstreamTransportHTTPSSE, compactPath: true, body: []byte(`{"tools":[{"type":"namespace","name":"mcp__codex_app","tools":[]}]}`), want: false},
|
||||
{name: "setup_token_keeps", account: setupToken, transport: OpenAIUpstreamTransportHTTPSSE, want: true},
|
||||
{name: "nil_account", account: nil, transport: OpenAIUpstreamTransportHTTPSSE, want: false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
require.Equal(t, tt.want, shouldKeepOpenAIResponsesToolCallNamespaces(
|
||||
tt.account, tt.transport, tt.passthroughEnabled, tt.compactPath,
|
||||
tt.account, tt.transport, tt.passthroughEnabled, tt.compactPath, tt.body,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -538,6 +538,34 @@ func TestOpenAIGatewayService_APIKeyStripsAllIndexedNamespacesBeforeFirstForward
|
||||
require.False(t, gjson.GetBytes(upstream.bodies[0], "input.1.namespace").Exists())
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceProactivelyStripsCrossProviderReasoningContent(t *testing.T) {
|
||||
body := []byte(`{"model":"gpt-5.5","stream":false,"store":true,"input":[` +
|
||||
`{"type":"message","role":"user","content":"one"},` +
|
||||
`{"type":"message","role":"assistant","content":"two"},` +
|
||||
`{"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"},` +
|
||||
`{"type":"function_call_output","call_id":"call_1","output":"ok"},` +
|
||||
`{"type":"message","role":"user","content":"five"},` +
|
||||
`{"type":"reasoning","summary":[{"type":"summary_text","text":"keep"}],"content":[{"type":"reasoning_text","text":"remove"}]}` +
|
||||
`]}`)
|
||||
upstream := &httpUpstreamRecorder{responses: []*http.Response{
|
||||
newOpenAIRejectedFieldTestResponse(http.StatusOK, `{"output":[],"usage":{"input_tokens":1,"output_tokens":1,"input_tokens_details":{"cached_tokens":0}}}`),
|
||||
}}
|
||||
|
||||
result, err := newOpenAIRejectedFieldTestService(upstream).Forward(
|
||||
context.Background(),
|
||||
newOpenAIRejectedFieldTestContext(body),
|
||||
newOpenAIRejectedFieldTestAccount(),
|
||||
body,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Len(t, upstream.bodies, 1, "reasoning content should be normalized before the first upstream request")
|
||||
require.Equal(t, "reasoning", gjson.GetBytes(upstream.bodies[0], "input.5.type").String())
|
||||
require.False(t, gjson.GetBytes(upstream.bodies[0], "input.5.content").Exists())
|
||||
require.Equal(t, "keep", gjson.GetBytes(upstream.bodies[0], "input.5.summary.0.text").String())
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OpenAIHTTPStripsInputNamespacesBeforeFirstForward(t *testing.T) {
|
||||
accounts := []struct {
|
||||
name string
|
||||
|
||||
@@ -12,8 +12,17 @@ const (
|
||||
// when an upstream HTTP/2 response stream is reset after the request started.
|
||||
OpenAIUpstreamHTTP2StreamErrorCode = "upstream_http2_stream_error"
|
||||
OpenAIUpstreamStreamReadErrorCode = "upstream_stream_read_error"
|
||||
// OpenAIUpstreamStreamTruncatedCode is returned when an upstream SSE stream
|
||||
// closes *cleanly* before delivering any terminal signal. A clean EOF carries
|
||||
// no transport error, so without this classification a truncated generation is
|
||||
// indistinguishable from a successful one.
|
||||
OpenAIUpstreamStreamTruncatedCode = "upstream_stream_truncated"
|
||||
)
|
||||
|
||||
// ErrOpenAIUpstreamStreamTruncated marks an upstream SSE stream that ended at
|
||||
// EOF — without a read error — before any terminal signal arrived.
|
||||
var ErrOpenAIUpstreamStreamTruncated = errors.New("upstream stream ended before any terminal chunk")
|
||||
|
||||
type openAIUpstreamStreamReadError struct {
|
||||
cause error
|
||||
clientCode string
|
||||
@@ -61,6 +70,9 @@ func OpenAIUpstreamStreamReadErrorDetails(err error) (code, message string, ok b
|
||||
|
||||
func classifyOpenAIUpstreamStreamReadError(err error) (code, message string) {
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrOpenAIUpstreamStreamTruncated) {
|
||||
return OpenAIUpstreamStreamTruncatedCode, "Upstream response stream ended before completion"
|
||||
}
|
||||
lower := strings.ToLower(err.Error())
|
||||
// net/http's HTTP/2 stream error is unexported. Its stable text contains
|
||||
// "stream error: stream ID ..."; match only the transport signature and
|
||||
|
||||
@@ -1248,6 +1248,10 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
// Match the handler close path and stay within the WebSocket control
|
||||
// frame limit; an oversized reason makes coder/websocket skip the
|
||||
// close frame, leaving the client with EOF instead of the status code.
|
||||
reason = truncateString(reason, 120)
|
||||
_ = clientConn.Close(status, reason)
|
||||
_ = clientConn.CloseNow()
|
||||
},
|
||||
@@ -1259,6 +1263,9 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
|
||||
if eventType == "response.created" {
|
||||
failureAccountSideEffectsApplied = false
|
||||
}
|
||||
if (eventType == "error" || eventType == "response.failed") && markOpenAIWSV2PassthroughCyberPolicy(c, payload) {
|
||||
return nil
|
||||
}
|
||||
errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(payload)
|
||||
isPreOutputRateLimit := eventType == "error" && !wroteDownstream && isOpenAIWSRateLimitError(errCodeRaw, errTypeRaw, errMsgRaw)
|
||||
if (eventType == "error" || eventType == "response.failed") && !failureAccountSideEffectsApplied && !isPreOutputRateLimit {
|
||||
@@ -1449,6 +1456,24 @@ func openAIWSPassthroughRelayClientClose(exit openaiwsv2.RelayExit, completedTur
|
||||
return 0, "", false
|
||||
}
|
||||
|
||||
func markOpenAIWSV2PassthroughCyberPolicy(c *gin.Context, payload []byte) bool {
|
||||
hit, code, message := detectOpenAICyberPolicy(payload)
|
||||
if !hit {
|
||||
return false
|
||||
}
|
||||
usage := OpenAIUsage{}
|
||||
parseOpenAIWSResponseUsageFromCompletedEvent(payload, &usage)
|
||||
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||
Code: code,
|
||||
Message: message,
|
||||
Body: truncateString(string(payload), 4096),
|
||||
UpstreamStatus: http.StatusOK,
|
||||
UpstreamInTok: usage.InputTokens,
|
||||
UpstreamOutTok: usage.OutputTokens,
|
||||
})
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) mapOpenAIWSPassthroughDialError(
|
||||
err error,
|
||||
statusCode int,
|
||||
|
||||
@@ -8,8 +8,10 @@ import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
coderws "github.com/coder/websocket"
|
||||
@@ -21,6 +23,7 @@ import (
|
||||
type stagedPassthroughFrame struct {
|
||||
messageType coderws.MessageType
|
||||
payload []byte
|
||||
err error
|
||||
}
|
||||
|
||||
type stagedPassthroughConn struct {
|
||||
@@ -42,6 +45,10 @@ func (c *stagedPassthroughConn) Send(payload string) {
|
||||
c.frames <- stagedPassthroughFrame{messageType: coderws.MessageText, payload: []byte(payload)}
|
||||
}
|
||||
|
||||
func (c *stagedPassthroughConn) Fail(err error) {
|
||||
c.frames <- stagedPassthroughFrame{err: err}
|
||||
}
|
||||
|
||||
func (c *stagedPassthroughConn) WriteJSON(context.Context, any) error { return nil }
|
||||
|
||||
func (c *stagedPassthroughConn) ReadMessage(ctx context.Context) ([]byte, error) {
|
||||
@@ -61,7 +68,7 @@ func (c *stagedPassthroughConn) ReadFrame(ctx context.Context) (coderws.MessageT
|
||||
case <-c.closed:
|
||||
return coderws.MessageText, nil, errOpenAIWSConnClosed
|
||||
case frame := <-c.frames:
|
||||
return frame.messageType, append([]byte(nil), frame.payload...), nil
|
||||
return frame.messageType, append([]byte(nil), frame.payload...), frame.err
|
||||
}
|
||||
}
|
||||
|
||||
@@ -152,6 +159,16 @@ func startPassthroughLifecycleServer(
|
||||
controlCtx context.Context,
|
||||
svc *OpenAIGatewayService,
|
||||
account *Account,
|
||||
) (*httptest.Server, <-chan error) {
|
||||
return startPassthroughLifecycleServerWithHooks(t, controlCtx, svc, account, nil)
|
||||
}
|
||||
|
||||
func startPassthroughLifecycleServerWithHooks(
|
||||
t *testing.T,
|
||||
controlCtx context.Context,
|
||||
svc *OpenAIGatewayService,
|
||||
account *Account,
|
||||
hooksFactory func(*gin.Context) *OpenAIWSIngressHooks,
|
||||
) (*httptest.Server, <-chan error) {
|
||||
t.Helper()
|
||||
serverErr := make(chan error, 1)
|
||||
@@ -184,11 +201,217 @@ func startPassthroughLifecycleServer(
|
||||
req := r.Clone(controlCtx)
|
||||
req.Header = req.Header.Clone()
|
||||
ginCtx.Request = req
|
||||
serverErr <- svc.ProxyResponsesWebSocketFromClient(controlCtx, ginCtx, conn, account, "sk-test", firstMessage, nil)
|
||||
var hooks *OpenAIWSIngressHooks
|
||||
if hooksFactory != nil {
|
||||
hooks = hooksFactory(ginCtx)
|
||||
}
|
||||
serverErr <- svc.ProxyResponsesWebSocketFromClient(controlCtx, ginCtx, conn, account, "sk-test", firstMessage, hooks)
|
||||
}))
|
||||
return server, serverErr
|
||||
}
|
||||
|
||||
func TestPassthroughLifecycle_CyberTerminalEventsMarkBeforeAfterTurn(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
events []string
|
||||
wantBody string
|
||||
wantMessage string
|
||||
wantInput int
|
||||
wantOutput int
|
||||
}{
|
||||
{
|
||||
name: "error",
|
||||
events: []string{
|
||||
`{"type":"error","error":{"code":"cyber_policy","message":"blocked by error event"},"usage":{"input_tokens":5,"output_tokens":1}}`,
|
||||
`{"type":"response.failed","response":{"id":"resp_error","error":{"code":"cyber_policy","message":"blocked by paired failed event"},"usage":{"input_tokens":9,"output_tokens":2}}}`,
|
||||
},
|
||||
wantBody: `"type":"error"`,
|
||||
wantMessage: "blocked by error event",
|
||||
wantInput: 5,
|
||||
wantOutput: 1,
|
||||
},
|
||||
{
|
||||
name: "response_failed",
|
||||
events: []string{
|
||||
`{"type":"response.failed","response":{"id":"resp_failed","error":{"code":"cyber_policy","message":"blocked by failed event"},"usage":{"input_tokens":9,"output_tokens":2}}}`,
|
||||
},
|
||||
wantBody: `"type":"response.failed"`,
|
||||
wantMessage: "blocked by failed event",
|
||||
wantInput: 9,
|
||||
wantOutput: 2,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
controlCtx, cancelControl := context.WithCancelCause(context.Background())
|
||||
defer cancelControl(context.Canceled)
|
||||
upstream := newStagedPassthroughConn()
|
||||
for _, event := range tt.events {
|
||||
upstream.Send(event)
|
||||
}
|
||||
|
||||
markSeen := make(chan CyberPolicyMark, 1)
|
||||
afterTurnCalls := atomic.Int32{}
|
||||
server, serverErr := startPassthroughLifecycleServerWithHooks(
|
||||
t,
|
||||
controlCtx,
|
||||
newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream),
|
||||
passthroughLifecycleAccount(),
|
||||
func(c *gin.Context) *OpenAIWSIngressHooks {
|
||||
return &OpenAIWSIngressHooks{AfterTurn: func(_ int, _ *OpenAIForwardResult, _ error) {
|
||||
afterTurnCalls.Add(1)
|
||||
if mark := GetOpsCyberPolicy(c); mark != nil {
|
||||
select {
|
||||
case markSeen <- *mark:
|
||||
default:
|
||||
}
|
||||
}
|
||||
}}
|
||||
},
|
||||
)
|
||||
defer server.Close()
|
||||
clientConn := dialPassthroughLifecycleClient(t, server)
|
||||
defer func() { _ = clientConn.CloseNow() }()
|
||||
|
||||
for range tt.events {
|
||||
_, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
select {
|
||||
case mark := <-markSeen:
|
||||
require.Equal(t, "cyber_policy", mark.Code)
|
||||
require.Equal(t, tt.wantMessage, mark.Message)
|
||||
require.Contains(t, mark.Body, tt.wantBody)
|
||||
require.Equal(t, http.StatusOK, mark.UpstreamStatus)
|
||||
require.Equal(t, tt.wantInput, mark.UpstreamInTok)
|
||||
require.Equal(t, tt.wantOutput, mark.UpstreamOutTok)
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("cyber mark was not visible to AfterTurn")
|
||||
}
|
||||
require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
|
||||
select {
|
||||
case <-serverErr:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("cyber passthrough test did not exit")
|
||||
}
|
||||
require.Equal(t, int32(1), afterTurnCalls.Load(), "error/response.failed pair must complete and record once")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPassthroughLifecycle_NonCyberFailureKeepsAccountSideEffects(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
controlCtx, cancelControl := context.WithCancelCause(context.Background())
|
||||
defer cancelControl(context.Canceled)
|
||||
upstream := newStagedPassthroughConn()
|
||||
upstream.Send(`{"type":"response.failed","response":{"id":"resp_non_cyber","error":{"type":"authentication_error","code":"invalid_api_key","status_code":401,"message":"credential rejected"},"usage":{"input_tokens":3,"output_tokens":1}}}`)
|
||||
repo := &openAIStream403AccountRepo{}
|
||||
svc := newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream)
|
||||
svc.rateLimitService = NewRateLimitService(repo, nil, svc.cfg, nil, nil)
|
||||
account := passthroughLifecycleAccount()
|
||||
|
||||
markSeen := make(chan *CyberPolicyMark, 1)
|
||||
server, serverErr := startPassthroughLifecycleServerWithHooks(
|
||||
t,
|
||||
controlCtx,
|
||||
svc,
|
||||
account,
|
||||
func(c *gin.Context) *OpenAIWSIngressHooks {
|
||||
return &OpenAIWSIngressHooks{AfterTurn: func(_ int, _ *OpenAIForwardResult, _ error) {
|
||||
markSeen <- GetOpsCyberPolicy(c)
|
||||
}}
|
||||
},
|
||||
)
|
||||
defer server.Close()
|
||||
clientConn := dialPassthroughLifecycleClient(t, server)
|
||||
defer func() { _ = clientConn.CloseNow() }()
|
||||
|
||||
event, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "response.failed", gjson.GetBytes(event, "type").String())
|
||||
select {
|
||||
case mark := <-markSeen:
|
||||
require.Nil(t, mark)
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("non-cyber terminal event did not complete its turn")
|
||||
}
|
||||
require.Equal(t, 1, repo.setErrorCalls, "non-cyber credential failure must retain account failure side effects")
|
||||
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
|
||||
require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
|
||||
select {
|
||||
case <-serverErr:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("non-cyber passthrough test did not exit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPassthroughLifecycle_CyberSkipsFailureAccountSideEffects(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
controlCtx, cancelControl := context.WithCancelCause(context.Background())
|
||||
defer cancelControl(context.Canceled)
|
||||
upstream := newStagedPassthroughConn()
|
||||
upstream.Send(`{"type":"response.failed","response":{"id":"resp_cyber_auth","error":{"type":"authentication_error","code":"cyber_policy","status_code":401,"message":"request blocked"}}}`)
|
||||
repo := &openAIStream403AccountRepo{}
|
||||
svc := newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream)
|
||||
svc.rateLimitService = NewRateLimitService(repo, nil, svc.cfg, nil, nil)
|
||||
account := passthroughLifecycleAccount()
|
||||
|
||||
server, serverErr := startPassthroughLifecycleServer(t, controlCtx, svc, account)
|
||||
defer server.Close()
|
||||
clientConn := dialPassthroughLifecycleClient(t, server)
|
||||
defer func() { _ = clientConn.CloseNow() }()
|
||||
|
||||
event, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "response.failed", gjson.GetBytes(event, "type").String())
|
||||
require.Zero(t, repo.setErrorCalls, "cyber_policy is request-scoped and must not cool down the account")
|
||||
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
|
||||
|
||||
require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
|
||||
select {
|
||||
case <-serverErr:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("cyber side-effect test did not exit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPassthroughLifecycle_CloseReasonTruncationPreservesUTF8(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
controlCtx, cancelControl := context.WithCancelCause(context.Background())
|
||||
defer cancelControl(context.Canceled)
|
||||
upstream := newStagedPassthroughConn()
|
||||
originalReason := strings.Repeat("a", 119) + "界"
|
||||
upstream.Fail(NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, originalReason, errors.New("policy rejected")))
|
||||
|
||||
server, serverErr := startPassthroughLifecycleServer(
|
||||
t,
|
||||
controlCtx,
|
||||
newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream),
|
||||
passthroughLifecycleAccount(),
|
||||
)
|
||||
defer server.Close()
|
||||
clientConn := dialPassthroughLifecycleClient(t, server)
|
||||
defer func() { _ = clientConn.CloseNow() }()
|
||||
|
||||
_, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second)
|
||||
var closeErr coderws.CloseError
|
||||
require.ErrorAs(t, err, &closeErr)
|
||||
require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code)
|
||||
require.True(t, utf8.ValidString(closeErr.Reason))
|
||||
require.LessOrEqual(t, len(closeErr.Reason), 120)
|
||||
require.Equal(t, strings.Repeat("a", 119), closeErr.Reason)
|
||||
|
||||
select {
|
||||
case <-serverErr:
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("passthrough close reason test did not exit")
|
||||
}
|
||||
}
|
||||
|
||||
func dialPassthroughLifecycleClient(t *testing.T, server *httptest.Server) *coderws.Conn {
|
||||
t.Helper()
|
||||
return dialPassthroughLifecycleClientWithPayload(t, server, `{"type":"response.create","model":"gpt-5.1","stream":false}`)
|
||||
|
||||
@@ -75,6 +75,8 @@ const (
|
||||
const (
|
||||
openAIImageRateLimitDefaultCooldown = time.Minute
|
||||
openAIImageRateLimitReason = "openai_image_rate_limited"
|
||||
openAIImageCapabilityLossCooldown = 30 * time.Minute
|
||||
openAIImageCapabilityLossReason = "openai_image_capability_lost"
|
||||
)
|
||||
|
||||
var openAIImageTryAgainPattern = regexp.MustCompile(`(?i)try again in\s+([0-9]+(?:\.[0-9]+)?)\s*(ms|s|sec|secs|second|seconds|m|min|mins|minute|minutes)`)
|
||||
@@ -2190,6 +2192,44 @@ func (s *RateLimitService) HandleOpenAIImageRateLimit(ctx context.Context, accou
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *RateLimitService) HandleOpenAIImageCapabilityLoss(ctx context.Context, account *Account, statusCode int, responseBody []byte) bool {
|
||||
if s == nil || account == nil || s.accountRepo == nil {
|
||||
return false
|
||||
}
|
||||
if account.Platform != PlatformOpenAI {
|
||||
return false
|
||||
}
|
||||
if !account.ShouldHandleErrorCode(statusCode) {
|
||||
slog.Info("openai_image_capability_loss_skipped_by_error_code_policy", "account_id", account.ID, "status_code", statusCode)
|
||||
return false
|
||||
}
|
||||
if !isOpenAIImageCapabilityLossError(statusCode, responseBody) {
|
||||
return false
|
||||
}
|
||||
|
||||
resetAt := time.Now().Add(openAIImageCapabilityLossCooldown)
|
||||
if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, openAIImageGenerationRateLimitKey, resetAt, openAIImageCapabilityLossReason); err != nil {
|
||||
slog.Warn("openai_image_capability_loss_set_model_rate_limit_failed", "account_id", account.ID, "scope", openAIImageGenerationRateLimitKey, "error", err)
|
||||
return true
|
||||
}
|
||||
slog.Info("openai_image_capability_lost", "account_id", account.ID, "scope", openAIImageGenerationRateLimitKey, "reset_at", resetAt, "reset_in", time.Until(resetAt).Truncate(time.Second))
|
||||
return true
|
||||
}
|
||||
|
||||
// isOpenAIImageCapabilityLossError reports whether upstream rejected the
|
||||
// image_generation tool choice that sub2api itself put into the request body.
|
||||
// Only meaningful for self-built images requests, where tools always carries a
|
||||
// matching image_generation entry — upstream saying otherwise means the account
|
||||
// lost the capability.
|
||||
func isOpenAIImageCapabilityLossError(statusCode int, body []byte) bool {
|
||||
if statusCode != http.StatusBadRequest || len(body) == 0 {
|
||||
return false
|
||||
}
|
||||
lower := strings.ToLower(string(body))
|
||||
return strings.Contains(lower, "image_generation") &&
|
||||
strings.Contains(lower, "not found in 'tools' parameter")
|
||||
}
|
||||
|
||||
func isOpenAIImageRateLimitError(statusCode int, body []byte) bool {
|
||||
if statusCode != http.StatusTooManyRequests || len(body) == 0 {
|
||||
return false
|
||||
|
||||
@@ -120,7 +120,11 @@ func TestOpenAIGatewayServiceForwardImages_ImageRateLimitReturnsFailoverAndCools
|
||||
require.Equal(t, openAIImageGenerationRateLimitKey, repo.modelRateLimitCalls[0].scope)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *testing.T) {
|
||||
// issue #6171:上游"回文字没回图"是**这一轮**的结果(模型选择了说话),不是账号能力
|
||||
// 失效。它同时被判为可重试(502)并驱动 failover,若还写 30 分钟账号级冷却,一次闲聊
|
||||
// 回复就会沿号池把每个被重试到的账号依次冷却掉。冷却仍保留给结构化上游证据,见
|
||||
// TestOpenAIGatewayServiceForwardImages_StructuredUnavailableCoolsImageCapability。
|
||||
func TestOpenAIGatewayServiceForwardImages_TextFallbackDoesNotCoolImageCapability(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repo := &modelNotFoundAccountRepoStub{}
|
||||
body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat"}`)
|
||||
@@ -154,7 +158,6 @@ func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *t
|
||||
},
|
||||
}
|
||||
|
||||
before := time.Now()
|
||||
result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "")
|
||||
|
||||
require.Nil(t, result)
|
||||
@@ -162,6 +165,56 @@ func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *t
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.False(t, failoverErr.RetryableOnSameAccount)
|
||||
// 换号行为不变:该判据仍足以放弃本账号重试这一次请求……
|
||||
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
|
||||
// ……但不再写任何账号级状态,否则重试会把冷却一路刷到整个号池。
|
||||
require.Empty(t, repo.modelRateLimitCalls,
|
||||
"模型回文字只说明这一轮没出图,不构成账号 30 分钟不可用的证据")
|
||||
}
|
||||
|
||||
// 对照不变式:上游 error 帧点名 image_generation_unavailable 时仍写冷却,
|
||||
// 保证 #6171 的修复没有把这项能力保护整个废掉。
|
||||
func TestOpenAIGatewayServiceForwardImages_StructuredUnavailableCoolsImageCapability(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repo := &modelNotFoundAccountRepoStub{}
|
||||
body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat"}`)
|
||||
upstreamSSE := "data: {\"type\":\"response.failed\",\"response\":{\"id\":\"r\",\"error\":" +
|
||||
"{\"type\":\"upstream_error\",\"code\":\"image_generation_unavailable\"," +
|
||||
"\"message\":\"image generation tool is not available for this account\"}}}\n\n"
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = req
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: repo,
|
||||
httpUpstream: &httpUpstreamRecorder{
|
||||
resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamSSE)),
|
||||
},
|
||||
},
|
||||
}
|
||||
parsed, err := svc.ParseOpenAIImagesRequest(c, body)
|
||||
require.NoError(t, err)
|
||||
account := &Account{
|
||||
ID: 206,
|
||||
Name: "openai-oauth",
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "token-123",
|
||||
},
|
||||
}
|
||||
|
||||
before := time.Now()
|
||||
result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "")
|
||||
|
||||
require.Nil(t, result)
|
||||
require.Error(t, err)
|
||||
require.Len(t, repo.modelRateLimitCalls, 1)
|
||||
call := repo.modelRateLimitCalls[0]
|
||||
require.Equal(t, account.ID, call.accountID)
|
||||
@@ -169,3 +222,120 @@ func TestOpenAIGatewayServiceForwardImages_TextFallbackCoolsImageCapability(t *t
|
||||
require.Equal(t, openAIImagesOAuthUnavailableReason, call.reason)
|
||||
require.WithinDuration(t, before.Add(openAIImagesOAuthUnavailableCooldown), call.resetAt, time.Second)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceForwardImages_CapabilityLossCoolsImageScope(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repo := &modelNotFoundAccountRepoStub{}
|
||||
body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat"}`)
|
||||
errorBody := `{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}`
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = req
|
||||
|
||||
svc := &OpenAIGatewayService{
|
||||
rateLimitService: &RateLimitService{accountRepo: repo},
|
||||
httpUpstream: &httpUpstreamRecorder{
|
||||
resp: &http.Response{
|
||||
StatusCode: http.StatusBadRequest,
|
||||
Header: http.Header{"X-Request-Id": []string{"req_img_capability_lost"}},
|
||||
Body: io.NopCloser(strings.NewReader(errorBody)),
|
||||
},
|
||||
},
|
||||
}
|
||||
parsed, err := svc.ParseOpenAIImagesRequest(c, body)
|
||||
require.NoError(t, err)
|
||||
account := &Account{
|
||||
ID: 205,
|
||||
Name: "openai-oauth",
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "token-123",
|
||||
},
|
||||
}
|
||||
|
||||
before := time.Now()
|
||||
result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "")
|
||||
|
||||
require.Nil(t, result)
|
||||
require.Error(t, err)
|
||||
require.Len(t, repo.modelRateLimitCalls, 1)
|
||||
call := repo.modelRateLimitCalls[0]
|
||||
require.Equal(t, account.ID, call.accountID)
|
||||
require.Equal(t, openAIImageGenerationRateLimitKey, call.scope)
|
||||
require.Equal(t, openAIImageCapabilityLossReason, call.reason)
|
||||
require.WithinDuration(t, before.Add(openAIImageCapabilityLossCooldown), call.resetAt, time.Second)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceHandleUpstreamError_PassthroughCapabilityLossDoesNotCool(t *testing.T) {
|
||||
repo := &modelNotFoundAccountRepoStub{}
|
||||
svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}}
|
||||
account := &Account{ID: 206, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
|
||||
body := []byte(`{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}`)
|
||||
|
||||
disabled := svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusBadRequest, http.Header{}, body, "gpt-5.5")
|
||||
|
||||
require.False(t, disabled)
|
||||
require.Empty(t, repo.modelRateLimitCalls)
|
||||
_, wholeAccountBlocked := svc.openaiAccountRuntimeBlockUntil.Load(account.ID)
|
||||
require.False(t, wholeAccountBlocked)
|
||||
}
|
||||
|
||||
func TestRateLimitServiceHandleOpenAIImageCapabilityLoss_IgnoresGenericBadRequest(t *testing.T) {
|
||||
repo := &modelNotFoundAccountRepoStub{}
|
||||
svc := &RateLimitService{accountRepo: repo}
|
||||
account := &Account{ID: 207, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
|
||||
body := []byte(`{"error":{"message":"Invalid type for input[0].arguments"}}`)
|
||||
|
||||
handled := svc.HandleOpenAIImageCapabilityLoss(context.Background(), account, http.StatusBadRequest, body)
|
||||
|
||||
require.False(t, handled)
|
||||
require.Empty(t, repo.modelRateLimitCalls)
|
||||
}
|
||||
|
||||
func TestRateLimitServiceHandleOpenAIImageCapabilityLoss_RespectsPlatformAndErrorCodePolicy(t *testing.T) {
|
||||
body := []byte(`{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}`)
|
||||
|
||||
t.Run("non_openai_platform", func(t *testing.T) {
|
||||
repo := &modelNotFoundAccountRepoStub{}
|
||||
svc := &RateLimitService{accountRepo: repo}
|
||||
account := &Account{ID: 208, Platform: PlatformAnthropic, Type: AccountTypeOAuth}
|
||||
|
||||
handled := svc.HandleOpenAIImageCapabilityLoss(context.Background(), account, http.StatusBadRequest, body)
|
||||
|
||||
require.False(t, handled)
|
||||
require.Empty(t, repo.modelRateLimitCalls)
|
||||
})
|
||||
|
||||
t.Run("custom_error_code_policy_excludes_400", func(t *testing.T) {
|
||||
repo := &modelNotFoundAccountRepoStub{}
|
||||
svc := &RateLimitService{accountRepo: repo}
|
||||
account := &Account{
|
||||
ID: 209,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Credentials: map[string]any{
|
||||
"custom_error_codes_enabled": true,
|
||||
"custom_error_codes": []any{float64(http.StatusTooManyRequests)},
|
||||
},
|
||||
}
|
||||
|
||||
require.False(t, account.ShouldHandleErrorCode(http.StatusBadRequest))
|
||||
handled := svc.HandleOpenAIImageCapabilityLoss(context.Background(), account, http.StatusBadRequest, body)
|
||||
|
||||
require.False(t, handled)
|
||||
require.Empty(t, repo.modelRateLimitCalls)
|
||||
})
|
||||
}
|
||||
|
||||
func TestIsOpenAIImageCapabilityLossError(t *testing.T) {
|
||||
capabilityLossBody := []byte(`{"error":{"message":"Tool choice 'image_generation' not found in 'tools' parameter.","param":"tool_choice","type":"invalid_request_error"}}`)
|
||||
genericBadRequestBody := []byte(`{"error":{"message":"Invalid type for input[0].arguments"}}`)
|
||||
|
||||
require.True(t, isOpenAIImageCapabilityLossError(http.StatusBadRequest, capabilityLossBody))
|
||||
require.False(t, isOpenAIImageCapabilityLossError(http.StatusBadRequest, genericBadRequestBody))
|
||||
require.False(t, isOpenAIImageCapabilityLossError(http.StatusTooManyRequests, capabilityLossBody))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
package testutil
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/repository"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
// NewRedisGatewayCache returns a real Redis-backed gateway cache for tests.
|
||||
func NewRedisGatewayCache(t *testing.T) service.GatewayCache {
|
||||
t.Helper()
|
||||
|
||||
redisServer := miniredis.RunT(t)
|
||||
redisClient := redis.NewClient(&redis.Options{Addr: redisServer.Addr()})
|
||||
t.Cleanup(func() { _ = redisClient.Close() })
|
||||
|
||||
return repository.NewGatewayCache(redisClient)
|
||||
}
|
||||
Reference in New Issue
Block a user