mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge pull request #3565 from zy6p/zy6p/pr-openai-ws-http-bridge
feat(openai-ws): 支持 http_bridge ingress 模式
This commit is contained in:
@@ -721,6 +721,9 @@ type GatewayConfig struct {
|
||||
// OpenAIPassthroughAllowTimeoutHeaders: OpenAI 透传模式是否放行客户端超时头
|
||||
// 关闭(默认)可避免 x-stainless-timeout 等头导致上游提前断流。
|
||||
OpenAIPassthroughAllowTimeoutHeaders bool `mapstructure:"openai_passthrough_allow_timeout_headers"`
|
||||
// OpenAICompactModel: /responses/compact 上游使用的模型。
|
||||
// compact 端点支持模型滞后于普通 /responses 时,可用该配置降级规避上游错误。
|
||||
OpenAICompactModel string `mapstructure:"openai_compact_model"`
|
||||
// OpenAIWS: OpenAI Responses WebSocket 配置(默认开启,可按需回滚到 HTTP)
|
||||
OpenAIWS GatewayOpenAIWSConfig `mapstructure:"openai_ws"`
|
||||
// OpenAIScheduler: OpenAI 高级调度器粘性逃逸配置
|
||||
@@ -867,7 +870,7 @@ func (c *UserMessageQueueConfig) GetEffectiveMode() string {
|
||||
type GatewayOpenAIWSConfig struct {
|
||||
// ModeRouterV2Enabled: 新版 WS mode 路由开关(默认 false;关闭时保持 legacy 行为)
|
||||
ModeRouterV2Enabled bool `mapstructure:"mode_router_v2_enabled"`
|
||||
// IngressModeDefault: ingress 默认模式(off/ctx_pool/passthrough)
|
||||
// IngressModeDefault: ingress 默认模式(off/ctx_pool/passthrough/http_bridge)
|
||||
IngressModeDefault string `mapstructure:"ingress_mode_default"`
|
||||
// Enabled: 全局总开关(默认 true)
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
@@ -1836,6 +1839,7 @@ func setDefaults() {
|
||||
viper.SetDefault("gateway.force_codex_cli", false)
|
||||
viper.SetDefault("gateway.codex_image_generation_bridge_enabled", false)
|
||||
viper.SetDefault("gateway.openai_passthrough_allow_timeout_headers", false)
|
||||
viper.SetDefault("gateway.openai_compact_model", "gpt-5.4")
|
||||
// OpenAI Responses WebSocket(默认开启;可通过 force_http 紧急回滚)
|
||||
viper.SetDefault("gateway.openai_ws.enabled", true)
|
||||
viper.SetDefault("gateway.openai_ws.mode_router_v2_enabled", false)
|
||||
@@ -2624,11 +2628,11 @@ func (c *Config) Validate() error {
|
||||
}
|
||||
if mode := strings.ToLower(strings.TrimSpace(c.Gateway.OpenAIWS.IngressModeDefault)); mode != "" {
|
||||
switch mode {
|
||||
case "off", "ctx_pool", "passthrough":
|
||||
case "off", "ctx_pool", "passthrough", "http_bridge":
|
||||
case "shared", "dedicated":
|
||||
slog.Warn("gateway.openai_ws.ingress_mode_default is deprecated, treating as ctx_pool; please update to off|ctx_pool|passthrough", "value", mode)
|
||||
slog.Warn("gateway.openai_ws.ingress_mode_default is deprecated, treating as ctx_pool; please update to off|ctx_pool|passthrough|http_bridge", "value", mode)
|
||||
default:
|
||||
return fmt.Errorf("gateway.openai_ws.ingress_mode_default must be one of off|ctx_pool|passthrough")
|
||||
return fmt.Errorf("gateway.openai_ws.ingress_mode_default must be one of off|ctx_pool|passthrough|http_bridge")
|
||||
}
|
||||
}
|
||||
if mode := strings.ToLower(strings.TrimSpace(c.Gateway.OpenAIWS.StoreDisabledConnMode)); mode != "" {
|
||||
|
||||
@@ -184,6 +184,23 @@ func TestLoadDefaultOpenAIWSConfig(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDefaultOpenAICompactModel(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
|
||||
cfg, err := Load()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "gpt-5.4", cfg.Gateway.OpenAICompactModel)
|
||||
}
|
||||
|
||||
func TestLoadOpenAICompactModelFromEnv(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
t.Setenv("GATEWAY_OPENAI_COMPACT_MODEL", "gpt-5.3-codex")
|
||||
|
||||
cfg, err := Load()
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "gpt-5.3-codex", cfg.Gateway.OpenAICompactModel)
|
||||
}
|
||||
|
||||
func TestLoadDefaultOpenAIHTTP2Enabled(t *testing.T) {
|
||||
resetViperWithJWTSecret(t)
|
||||
|
||||
@@ -1678,7 +1695,7 @@ func TestValidateConfig_OpenAIWSRules(t *testing.T) {
|
||||
wantErr: "gateway.openai_ws.store_disabled_conn_mode",
|
||||
},
|
||||
{
|
||||
name: "ingress_mode_default 必须为 off|ctx_pool|passthrough",
|
||||
name: "ingress_mode_default 必须为 off|ctx_pool|passthrough|http_bridge",
|
||||
mutate: func(c *Config) { c.Gateway.OpenAIWS.IngressModeDefault = "invalid" },
|
||||
wantErr: "gateway.openai_ws.ingress_mode_default",
|
||||
},
|
||||
|
||||
@@ -1349,7 +1349,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
|
||||
subscription, _ := middleware2.GetSubscriptionFromContext(c)
|
||||
requestPlatform := openAICompatibleRequestPlatform(apiKey)
|
||||
requiredTransport := service.OpenAIUpstreamTransportResponsesWebsocketV2
|
||||
requiredTransport := service.OpenAIUpstreamTransportResponsesWebsocketV2Ingress
|
||||
if requestPlatform == service.PlatformGrok {
|
||||
requiredTransport = service.OpenAIUpstreamTransportHTTPSSE
|
||||
}
|
||||
|
||||
@@ -1474,6 +1474,7 @@ const (
|
||||
OpenAIWSIngressModeDedicated = "dedicated"
|
||||
OpenAIWSIngressModeCtxPool = "ctx_pool"
|
||||
OpenAIWSIngressModePassthrough = "passthrough"
|
||||
OpenAIWSIngressModeHTTPBridge = "http_bridge"
|
||||
)
|
||||
|
||||
func normalizeOpenAIWSIngressMode(mode string) string {
|
||||
@@ -1484,6 +1485,8 @@ func normalizeOpenAIWSIngressMode(mode string) string {
|
||||
return OpenAIWSIngressModeCtxPool
|
||||
case OpenAIWSIngressModePassthrough:
|
||||
return OpenAIWSIngressModePassthrough
|
||||
case OpenAIWSIngressModeHTTPBridge:
|
||||
return OpenAIWSIngressModeHTTPBridge
|
||||
case OpenAIWSIngressModeShared:
|
||||
return OpenAIWSIngressModeShared
|
||||
case OpenAIWSIngressModeDedicated:
|
||||
|
||||
@@ -229,6 +229,17 @@ func TestAccount_ResolveOpenAIResponsesWebSocketV2Mode(t *testing.T) {
|
||||
require.Equal(t, OpenAIWSIngressModePassthrough, account.ResolveOpenAIResponsesWebSocketV2Mode(OpenAIWSIngressModeCtxPool))
|
||||
})
|
||||
|
||||
t.Run("oauth mode supports http_bridge", func(t *testing.T) {
|
||||
account := &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Extra: map[string]any{
|
||||
"openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModeHTTPBridge,
|
||||
},
|
||||
}
|
||||
require.Equal(t, OpenAIWSIngressModeHTTPBridge, account.ResolveOpenAIResponsesWebSocketV2Mode(OpenAIWSIngressModeCtxPool))
|
||||
})
|
||||
|
||||
t.Run("legacy enabled maps to ctx_pool", func(t *testing.T) {
|
||||
account := &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
|
||||
@@ -1420,6 +1420,18 @@ func (s *OpenAIGatewayService) isOpenAIAccountTransportCompatible(account *Accou
|
||||
if s == nil || account == nil {
|
||||
return false
|
||||
}
|
||||
if requiredTransport == OpenAIUpstreamTransportResponsesWebsocketV2Ingress {
|
||||
if s.cfg == nil || !s.cfg.Gateway.OpenAIWS.ModeRouterV2Enabled {
|
||||
return s.getOpenAIWSProtocolResolver().Resolve(account).Transport == OpenAIUpstreamTransportResponsesWebsocketV2
|
||||
}
|
||||
mode := account.ResolveOpenAIResponsesWebSocketV2Mode(s.cfg.Gateway.OpenAIWS.IngressModeDefault)
|
||||
switch mode {
|
||||
case OpenAIWSIngressModeCtxPool, OpenAIWSIngressModePassthrough, OpenAIWSIngressModeHTTPBridge, OpenAIWSIngressModeShared, OpenAIWSIngressModeDedicated:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
return s.getOpenAIWSProtocolResolver().Resolve(account).Transport == requiredTransport
|
||||
}
|
||||
|
||||
|
||||
@@ -2391,6 +2391,7 @@ func TestDefaultOpenAIAccountScheduler_IsAccountTransportCompatible_Branches(t *
|
||||
require.True(t, scheduler.isAccountTransportCompatible(nil, OpenAIUpstreamTransportAny))
|
||||
require.True(t, scheduler.isAccountTransportCompatible(nil, OpenAIUpstreamTransportHTTPSSE))
|
||||
require.False(t, scheduler.isAccountTransportCompatible(nil, OpenAIUpstreamTransportResponsesWebsocketV2))
|
||||
require.False(t, scheduler.isAccountTransportCompatible(nil, OpenAIUpstreamTransportResponsesWebsocketV2Ingress))
|
||||
|
||||
cfg := newSchedulerTestOpenAIWSV2Config()
|
||||
scheduler.service = &OpenAIGatewayService{cfg: cfg}
|
||||
@@ -2406,6 +2407,15 @@ func TestDefaultOpenAIAccountScheduler_IsAccountTransportCompatible_Branches(t *
|
||||
},
|
||||
}
|
||||
require.True(t, scheduler.isAccountTransportCompatible(account, OpenAIUpstreamTransportResponsesWebsocketV2))
|
||||
require.True(t, scheduler.isAccountTransportCompatible(account, OpenAIUpstreamTransportResponsesWebsocketV2Ingress))
|
||||
|
||||
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
|
||||
account.Extra["openai_apikey_responses_websockets_v2_mode"] = OpenAIWSIngressModeHTTPBridge
|
||||
require.False(t, scheduler.isAccountTransportCompatible(account, OpenAIUpstreamTransportResponsesWebsocketV2))
|
||||
require.True(t, scheduler.isAccountTransportCompatible(account, OpenAIUpstreamTransportResponsesWebsocketV2Ingress))
|
||||
|
||||
account.Extra["openai_apikey_responses_websockets_v2_mode"] = OpenAIWSIngressModeOff
|
||||
require.False(t, scheduler.isAccountTransportCompatible(account, OpenAIUpstreamTransportResponsesWebsocketV2Ingress))
|
||||
}
|
||||
|
||||
func int64PtrForTest(v int64) *int64 {
|
||||
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
@@ -35,6 +36,21 @@ type httpUpstreamRecorder struct {
|
||||
err error
|
||||
}
|
||||
|
||||
type passthroughErrReadCloser struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (r passthroughErrReadCloser) Read(_ []byte) (int, error) {
|
||||
if r.err != nil {
|
||||
return 0, r.err
|
||||
}
|
||||
return 0, io.ErrUnexpectedEOF
|
||||
}
|
||||
|
||||
func (r passthroughErrReadCloser) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *httpUpstreamRecorder) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) {
|
||||
u.lastReq = req
|
||||
u.lastProxyURL = proxyURL
|
||||
@@ -799,7 +815,7 @@ func TestOpenAIGatewayService_OAuthPassthrough_UpstreamErrorIncludesPassthroughF
|
||||
require.Equal(t, "http_error", arr[len(arr)-1].Kind)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *testing.T) {
|
||||
func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
originalBody := []byte(`{"model":"gpt-5.2","stream":false,"instructions":"local-test-instructions","input":[{"type":"text","text":"hi"}]}`)
|
||||
|
||||
@@ -825,11 +841,12 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
}
|
||||
|
||||
testCases := []struct {
|
||||
name string
|
||||
accountType string
|
||||
statusCode int
|
||||
body string
|
||||
assertRepo func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time)
|
||||
name string
|
||||
accountType string
|
||||
statusCode int
|
||||
body string
|
||||
expectFailover bool
|
||||
assertRepo func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time)
|
||||
}{
|
||||
{
|
||||
name: "oauth_429_rate_limit",
|
||||
@@ -839,6 +856,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
resetAt := time.Now().Add(7 * 24 * time.Hour).Unix()
|
||||
return fmt.Sprintf(`{"error":{"message":"The usage limit has been reached","type":"usage_limit_reached","resets_at":%d}}`, resetAt)
|
||||
}(),
|
||||
expectFailover: true,
|
||||
assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) {
|
||||
require.Len(t, repo.rateLimitCalls, 1)
|
||||
require.Empty(t, repo.overloadCalls)
|
||||
@@ -846,16 +864,50 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "oauth_529_overload",
|
||||
accountType: AccountTypeOAuth,
|
||||
statusCode: 529,
|
||||
body: `{"error":{"message":"server overloaded","type":"server_error"}}`,
|
||||
name: "oauth_529_overload",
|
||||
accountType: AccountTypeOAuth,
|
||||
statusCode: 529,
|
||||
body: `{"error":{"message":"server overloaded","type":"server_error"}}`,
|
||||
expectFailover: true,
|
||||
assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) {
|
||||
require.Empty(t, repo.rateLimitCalls)
|
||||
require.Len(t, repo.overloadCalls, 1)
|
||||
require.WithinDuration(t, start.Add(10*time.Minute), repo.overloadCalls[0], 5*time.Second)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "oauth_502_bad_gateway",
|
||||
accountType: AccountTypeOAuth,
|
||||
statusCode: http.StatusBadGateway,
|
||||
body: `{"error":{"message":"bad gateway","type":"server_error"}}`,
|
||||
expectFailover: false,
|
||||
assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) {
|
||||
require.Empty(t, repo.rateLimitCalls)
|
||||
require.Empty(t, repo.overloadCalls)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "oauth_503_unavailable",
|
||||
accountType: AccountTypeOAuth,
|
||||
statusCode: http.StatusServiceUnavailable,
|
||||
body: `{"error":{"message":"service unavailable","type":"server_error"}}`,
|
||||
expectFailover: false,
|
||||
assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) {
|
||||
require.Empty(t, repo.rateLimitCalls)
|
||||
require.Empty(t, repo.overloadCalls)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "oauth_504_gateway_timeout",
|
||||
accountType: AccountTypeOAuth,
|
||||
statusCode: http.StatusGatewayTimeout,
|
||||
body: `{"error":{"message":"gateway timeout","type":"server_error"}}`,
|
||||
expectFailover: false,
|
||||
assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) {
|
||||
require.Empty(t, repo.rateLimitCalls)
|
||||
require.Empty(t, repo.overloadCalls)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "apikey_429_rate_limit",
|
||||
accountType: AccountTypeAPIKey,
|
||||
@@ -864,6 +916,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
resetAt := time.Now().Add(7 * 24 * time.Hour).Unix()
|
||||
return fmt.Sprintf(`{"error":{"message":"The usage limit has been reached","type":"usage_limit_reached","resets_at":%d}}`, resetAt)
|
||||
}(),
|
||||
expectFailover: true,
|
||||
assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) {
|
||||
require.Len(t, repo.rateLimitCalls, 1)
|
||||
require.Empty(t, repo.overloadCalls)
|
||||
@@ -871,10 +924,11 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "apikey_529_overload",
|
||||
accountType: AccountTypeAPIKey,
|
||||
statusCode: 529,
|
||||
body: `{"error":{"message":"server overloaded","type":"server_error"}}`,
|
||||
name: "apikey_529_overload",
|
||||
accountType: AccountTypeAPIKey,
|
||||
statusCode: 529,
|
||||
body: `{"error":{"message":"server overloaded","type":"server_error"}}`,
|
||||
expectFailover: true,
|
||||
assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) {
|
||||
require.Empty(t, repo.rateLimitCalls)
|
||||
require.Len(t, repo.overloadCalls, 1)
|
||||
@@ -919,9 +973,15 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
require.Error(t, err)
|
||||
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.Equal(t, tc.statusCode, failoverErr.StatusCode)
|
||||
require.False(t, c.Writer.Written(), "429/529 passthrough 应返回 failover 错误给上层换号,而不是直接向客户端写响应")
|
||||
if tc.expectFailover {
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.Equal(t, tc.statusCode, failoverErr.StatusCode)
|
||||
require.False(t, c.Writer.Written(), "retryable passthrough 错误应返回 failover 错误给上层换号,而不是直接向客户端写响应")
|
||||
} else {
|
||||
require.False(t, errors.As(err, &failoverErr))
|
||||
require.True(t, c.Writer.Written(), "非 failover 的 passthrough http 错误应直接写回客户端")
|
||||
require.Equal(t, tc.statusCode, rec.Code)
|
||||
}
|
||||
|
||||
v, ok := c.Get(OpsUpstreamErrorsKey)
|
||||
require.True(t, ok)
|
||||
@@ -929,7 +989,11 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
require.True(t, ok)
|
||||
require.NotEmpty(t, arr)
|
||||
require.True(t, arr[len(arr)-1].Passthrough)
|
||||
require.Equal(t, "failover", arr[len(arr)-1].Kind)
|
||||
if tc.expectFailover {
|
||||
require.Equal(t, "failover", arr[len(arr)-1].Kind)
|
||||
} else {
|
||||
require.Equal(t, "http_error", arr[len(arr)-1].Kind)
|
||||
}
|
||||
require.Equal(t, tc.statusCode, arr[len(arr)-1].UpstreamStatusCode)
|
||||
|
||||
tc.assertRepo(t, repo, start)
|
||||
@@ -937,6 +1001,73 @@ func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *test
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OpenAIPassthrough_CompactNetworkErrorsTriggerFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
resp *http.Response
|
||||
err error
|
||||
expectFailover bool
|
||||
}{
|
||||
{
|
||||
name: "request_error",
|
||||
err: errors.New("stream disconnected before completion"),
|
||||
expectFailover: true,
|
||||
},
|
||||
{
|
||||
name: "read_error",
|
||||
resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-compact"}},
|
||||
Body: passthroughErrReadCloser{err: io.ErrUnexpectedEOF},
|
||||
},
|
||||
expectFailover: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/compact", bytes.NewReader(nil))
|
||||
c.Request.Header.Set("User-Agent", "codex_cli_rs/0.1.0")
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: tt.resp, err: tt.err}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: &config.Config{Gateway: config.GatewayConfig{ForceCodexCLI: false}},
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
account := &Account{
|
||||
ID: 123,
|
||||
Name: "acc",
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc"},
|
||||
Extra: map[string]any{"openai_passthrough": true},
|
||||
Status: StatusActive,
|
||||
Schedulable: true,
|
||||
RateMultiplier: f64p(1),
|
||||
}
|
||||
body := []byte(`{"model":"gpt-5.5","instructions":"local-test-instructions","input":[{"type":"text","text":"compact me"}]}`)
|
||||
|
||||
_, err := svc.Forward(context.Background(), c, account, body)
|
||||
require.Error(t, err)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
if tt.expectFailover {
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
|
||||
require.False(t, c.Writer.Written(), "compact 网络错误应交给外层 failover,而不是直接写回客户端")
|
||||
} else {
|
||||
require.False(t, errors.As(err, &failoverErr))
|
||||
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
|
||||
require.False(t, c.Writer.Written())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OAuthPassthrough_NonCodexUAFallbackToCodexUA(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
@@ -2524,12 +2524,14 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
hooks,
|
||||
wsDecision,
|
||||
)
|
||||
case OpenAIWSIngressModeHTTPBridge:
|
||||
forceHTTPBridge = true
|
||||
case OpenAIWSIngressModeCtxPool, OpenAIWSIngressModeShared, OpenAIWSIngressModeDedicated:
|
||||
// continue
|
||||
default:
|
||||
return NewOpenAIWSClientCloseError(
|
||||
coderws.StatusPolicyViolation,
|
||||
"websocket mode only supports ctx_pool/passthrough",
|
||||
"websocket mode only supports ctx_pool/passthrough/http_bridge",
|
||||
nil,
|
||||
)
|
||||
}
|
||||
@@ -2840,7 +2842,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
||||
}
|
||||
refreshIngressRouteState(firstPayload)
|
||||
|
||||
if s.shouldBridgeOpenAIWSHTTP(account, firstPayload.payloadBytes, firstPayload.previousResponseID) {
|
||||
if forceHTTPBridge || s.shouldBridgeOpenAIWSHTTP(account, firstPayload.payloadBytes, firstPayload.previousResponseID) {
|
||||
logOpenAIWSModeInfo(
|
||||
"ingress_ws_http_bridge_start account_id=%d account_type=%s payload_bytes=%d threshold_bytes=%d has_session_hash=%v store_disabled=%v",
|
||||
account.ID,
|
||||
|
||||
@@ -830,6 +830,152 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_PassthroughHeade
|
||||
require.Equal(t, "turn-meta-1", captureDialer.lastHeaders.Get(openAIWSTurnMetadataHeader))
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_HTTPBridgeModeRelaysHTTPStream(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
cfg := &config.Config{}
|
||||
cfg.Security.URLAllowlist.Enabled = false
|
||||
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
|
||||
cfg.Gateway.OpenAIWS.Enabled = true
|
||||
cfg.Gateway.OpenAIWS.OAuthEnabled = true
|
||||
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
|
||||
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
|
||||
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
|
||||
cfg.Gateway.OpenAIWS.IngressModeDefault = OpenAIWSIngressModeCtxPool
|
||||
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
|
||||
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
|
||||
|
||||
upstream := &httpUpstreamRecorder{
|
||||
resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_bridge_1"}},
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"hello\"}\n\n" +
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_http_bridge_1\",\"usage\":{\"input_tokens\":2,\"output_tokens\":1,\"input_tokens_details\":{\"cached_tokens\":1}}}}\n\n" +
|
||||
"data: [DONE]\n\n",
|
||||
)),
|
||||
},
|
||||
}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: cfg,
|
||||
httpUpstream: upstream,
|
||||
cache: &stubGatewayCache{},
|
||||
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
|
||||
toolCorrector: NewCodexToolCorrector(),
|
||||
}
|
||||
|
||||
account := &Account{
|
||||
ID: 552,
|
||||
Name: "openai-ingress-http-bridge",
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Status: StatusActive,
|
||||
Schedulable: true,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "sk-test",
|
||||
},
|
||||
Extra: map[string]any{
|
||||
"openai_apikey_responses_websockets_v2_mode": OpenAIWSIngressModeHTTPBridge,
|
||||
},
|
||||
}
|
||||
|
||||
serverErrCh := make(chan error, 1)
|
||||
resultCh := make(chan *OpenAIForwardResult, 1)
|
||||
hooks := &OpenAIWSIngressHooks{
|
||||
AfterTurn: func(_ int, result *OpenAIForwardResult, turnErr error) {
|
||||
if turnErr == nil && result != nil {
|
||||
resultCh <- result
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
|
||||
CompressionMode: coderws.CompressionContextTakeover,
|
||||
})
|
||||
if err != nil {
|
||||
serverErrCh <- err
|
||||
return
|
||||
}
|
||||
defer func() {
|
||||
_ = conn.CloseNow()
|
||||
}()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
ginCtx, _ := gin.CreateTestContext(rec)
|
||||
req := r.Clone(r.Context())
|
||||
req.Header = req.Header.Clone()
|
||||
req.Header.Set("User-Agent", "unit-test-agent/1.0")
|
||||
ginCtx.Request = req
|
||||
|
||||
readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
|
||||
msgType, firstMessage, readErr := conn.Read(readCtx)
|
||||
cancel()
|
||||
if readErr != nil {
|
||||
serverErrCh <- readErr
|
||||
return
|
||||
}
|
||||
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
|
||||
serverErrCh <- errors.New("unsupported websocket client message type")
|
||||
return
|
||||
}
|
||||
|
||||
serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, hooks)
|
||||
}))
|
||||
defer wsServer.Close()
|
||||
|
||||
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
|
||||
cancelDial()
|
||||
require.NoError(t, err)
|
||||
defer func() {
|
||||
_ = clientConn.CloseNow()
|
||||
}()
|
||||
|
||||
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false}`))
|
||||
cancelWrite()
|
||||
require.NoError(t, err)
|
||||
|
||||
readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, event1, readErr1 := clientConn.Read(readCtx)
|
||||
cancelRead()
|
||||
require.NoError(t, readErr1)
|
||||
require.Equal(t, "response.output_text.delta", gjson.GetBytes(event1, "type").String())
|
||||
require.Equal(t, "hello", gjson.GetBytes(event1, "delta").String())
|
||||
|
||||
readCtx2, cancelRead2 := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
_, event2, readErr2 := clientConn.Read(readCtx2)
|
||||
cancelRead2()
|
||||
require.NoError(t, readErr2)
|
||||
require.Equal(t, "response.completed", gjson.GetBytes(event2, "type").String())
|
||||
require.Equal(t, "resp_http_bridge_1", gjson.GetBytes(event2, "response.id").String())
|
||||
|
||||
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
|
||||
|
||||
select {
|
||||
case serverErr := <-serverErrCh:
|
||||
require.NoError(t, serverErr)
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("等待 http_bridge websocket 结束超时")
|
||||
}
|
||||
|
||||
select {
|
||||
case result := <-resultCh:
|
||||
require.Equal(t, "resp_http_bridge_1", result.RequestID)
|
||||
require.True(t, result.OpenAIWSMode)
|
||||
require.Equal(t, 2, result.Usage.InputTokens)
|
||||
require.Equal(t, 1, result.Usage.OutputTokens)
|
||||
require.Equal(t, 1, result.Usage.CacheReadInputTokens)
|
||||
require.NotNil(t, result.FirstTokenMs)
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("未收到 http_bridge turn 结果回调")
|
||||
}
|
||||
|
||||
require.NotNil(t, upstream.lastReq, "http_bridge 模式应调用 HTTP 上游")
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_ModeOffReturnsPolicyViolation(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
@@ -10,6 +10,9 @@ const (
|
||||
OpenAIUpstreamTransportHTTPSSE OpenAIUpstreamTransport = "http_sse"
|
||||
OpenAIUpstreamTransportResponsesWebsocket OpenAIUpstreamTransport = "responses_websockets"
|
||||
OpenAIUpstreamTransportResponsesWebsocketV2 OpenAIUpstreamTransport = "responses_websockets_v2"
|
||||
// OpenAIUpstreamTransportResponsesWebsocketV2Ingress 用于 WS ingress 入口选账号:
|
||||
// mode_router_v2 开启时允许 ctx_pool/passthrough/http_bridge,拒绝 off。
|
||||
OpenAIUpstreamTransportResponsesWebsocketV2Ingress OpenAIUpstreamTransport = "responses_websockets_v2_ingress"
|
||||
)
|
||||
|
||||
// OpenAIWSProtocolDecision 表示协议决策结果。
|
||||
@@ -71,6 +74,8 @@ func (r *defaultOpenAIWSProtocolResolver) Resolve(account *Account) OpenAIWSProt
|
||||
return openAIWSHTTPDecision("account_mode_off")
|
||||
case OpenAIWSIngressModeCtxPool, OpenAIWSIngressModePassthrough:
|
||||
// continue
|
||||
case OpenAIWSIngressModeHTTPBridge:
|
||||
return openAIWSHTTPDecision("ws_v2_mode_http_bridge")
|
||||
case OpenAIWSIngressModeShared, OpenAIWSIngressModeDedicated:
|
||||
// 历史值兼容:按 ctx_pool 处理。
|
||||
mode = OpenAIWSIngressModeCtxPool
|
||||
|
||||
@@ -202,6 +202,20 @@ func TestOpenAIWSProtocolResolver_Resolve_ModeRouterV2(t *testing.T) {
|
||||
require.Equal(t, "ws_v2_mode_passthrough", decision.Reason)
|
||||
})
|
||||
|
||||
t.Run("http_bridge mode routes to http_sse", func(t *testing.T) {
|
||||
httpBridgeAccount := &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Extra: map[string]any{
|
||||
"openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModeHTTPBridge,
|
||||
},
|
||||
}
|
||||
decision := NewOpenAIWSProtocolResolver(cfg).Resolve(httpBridgeAccount)
|
||||
require.Equal(t, OpenAIUpstreamTransportHTTPSSE, decision.Transport)
|
||||
require.Equal(t, "ws_v2_mode_http_bridge", decision.Reason)
|
||||
})
|
||||
|
||||
t.Run("non-positive concurrency is rejected in v2 router", func(t *testing.T) {
|
||||
invalidConcurrency := &Account{
|
||||
Platform: PlatformOpenAI,
|
||||
|
||||
Reference in New Issue
Block a user