From f4e3eb1c5a23bc6b11896632d6bde913a20633bb Mon Sep 17 00:00:00 2001 From: Qi HU Date: Thu, 27 Aug 2026 16:59:15 +0800 Subject: [PATCH] fix(openai): detect cyber policy in ws v2 passthrough --- .../handler/openai_gateway_handler.go | 7 + .../openai_ws_v2_passthrough_cyber_test.go | 319 ++++++++++++++++++ .../openai_ws_v2_passthrough_adapter.go | 25 ++ ...openai_ws_v2_passthrough_lifecycle_test.go | 227 ++++++++++++- backend/internal/testutil/redis.go | 21 ++ 5 files changed, 597 insertions(+), 2 deletions(-) create mode 100644 backend/internal/handler/openai_ws_v2_passthrough_cyber_test.go create mode 100644 backend/internal/testutil/redis.go diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 32cea2f073..6b2fda70af 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -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 } diff --git a/backend/internal/handler/openai_ws_v2_passthrough_cyber_test.go b/backend/internal/handler/openai_ws_v2_passthrough_cyber_test.go new file mode 100644 index 0000000000..49e17c3e75 --- /dev/null +++ b/backend/internal/handler/openai_ws_v2_passthrough_cyber_test.go @@ -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") + } +} diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index 60e7b41416..6510a3731f 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -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, diff --git a/backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go b/backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go index ab7d704b04..5fa6a7def9 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go +++ b/backend/internal/service/openai_ws_v2_passthrough_lifecycle_test.go @@ -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}`) diff --git a/backend/internal/testutil/redis.go b/backend/internal/testutil/redis.go new file mode 100644 index 0000000000..caa74a411d --- /dev/null +++ b/backend/internal/testutil/redis.go @@ -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) +}