Merge pull request #6293 from specialpointcentral/fix/ws-v2-passthrough-cyber

fix(openai): detect cyber policy in ws v2 passthrough
This commit is contained in:
Wesley Liddick
2026-08-28 12:47:07 +08:00
committed by GitHub
5 changed files with 597 additions and 2 deletions
@@ -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")
}
}
@@ -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}`)
+21
View File
@@ -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)
}