Files
sub2api/backend/internal/service/openai_ws_client_read_test.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

144 lines
4.4 KiB
Go

package service
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
coderws "github.com/coder/websocket"
"github.com/stretchr/testify/require"
)
func TestReadOpenAIWSClientMessage_ControlCloseFrames(t *testing.T) {
tests := []struct {
name string
timeout time.Duration
timeoutStatus coderws.StatusCode
timeoutReason string
cancelCause error
wantStatus coderws.StatusCode
wantReason string
}{
{
name: "inter-turn idle sends normal close",
timeout: 25 * time.Millisecond,
timeoutStatus: coderws.StatusNormalClosure,
timeoutReason: "websocket idle timeout",
wantStatus: coderws.StatusNormalClosure,
wantReason: "websocket idle timeout",
},
{
name: "first message timeout sends policy close",
timeout: 25 * time.Millisecond,
timeoutStatus: coderws.StatusPolicyViolation,
timeoutReason: "missing first response.create message",
wantStatus: coderws.StatusPolicyViolation,
wantReason: "missing first response.create message",
},
{
name: "lease loss sends retry close",
cancelCause: ErrOpenAIWSIngressLeaseLost,
wantStatus: coderws.StatusTryAgainLater,
wantReason: "websocket ingress capacity lease lost; please reconnect",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
controlCtx, cancelControl := context.WithCancelCause(context.Background())
defer cancelControl(context.Canceled)
serverResult := make(chan error, 1)
readStarted := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, nil)
if err != nil {
serverResult <- err
return
}
defer func() { _ = conn.CloseNow() }()
close(readStarted)
_, _, err = ReadOpenAIWSClientMessage(
controlCtx,
conn,
tt.timeout,
tt.timeoutStatus,
tt.timeoutReason,
)
serverResult <- err
}))
defer server.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(server.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
<-readStarted
if tt.cancelCause != nil {
cancelControl(tt.cancelCause)
}
readCtx, cancelRead := context.WithTimeout(context.Background(), time.Second)
_, _, err = clientConn.Read(readCtx)
cancelRead()
var clientClose coderws.CloseError
require.ErrorAs(t, err, &clientClose)
require.Equal(t, tt.wantStatus, clientClose.Code)
require.Equal(t, tt.wantReason, clientClose.Reason)
select {
case serverErr := <-serverResult:
var closeErr *OpenAIWSClientCloseError
require.ErrorAs(t, serverErr, &closeErr)
require.Equal(t, tt.wantStatus, closeErr.StatusCode())
require.Equal(t, tt.wantReason, closeErr.Reason())
case <-time.After(time.Second):
t.Fatal("server read goroutine did not exit after close handshake")
}
})
}
}
func TestReadOpenAIWSClientMessage_ParentCancellationStillJoinsRead(t *testing.T) {
controlCtx, cancelControl := context.WithCancelCause(context.Background())
serverResult := make(chan error, 1)
readStarted := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, nil)
if err != nil {
serverResult <- err
return
}
defer func() { _ = conn.CloseNow() }()
close(readStarted)
_, _, err = ReadOpenAIWSClientMessage(controlCtx, conn, 0, 0, "")
serverResult <- err
}))
defer server.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), time.Second)
clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(server.URL, "http"), nil)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
<-readStarted
cancelControl(errors.New("server shutting down"))
readCtx, cancelRead := context.WithTimeout(context.Background(), time.Second)
_, _, err = clientConn.Read(readCtx)
cancelRead()
var clientClose coderws.CloseError
require.ErrorAs(t, err, &clientClose)
require.Equal(t, coderws.StatusGoingAway, clientClose.Code)
require.Equal(t, "websocket request canceled", clientClose.Reason)
select {
case <-serverResult:
case <-time.After(time.Second):
t.Fatal("server read goroutine leaked after parent cancellation")
}
}