mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
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:
@@ -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}`)
|
||||
|
||||
@@ -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