Files
sub2api/backend/internal/service/openai_ws_client_read.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

118 lines
2.9 KiB
Go

package service
import (
"context"
"errors"
"time"
coderws "github.com/coder/websocket"
)
type openAIWSClientReadResult struct {
messageType coderws.MessageType
payload []byte
err error
}
// ReadOpenAIWSClientMessage keeps one reader alive while control events send
// their close frame, then closes the transport and joins that reader.
func ReadOpenAIWSClientMessage(
controlCtx context.Context,
conn *coderws.Conn,
timeout time.Duration,
timeoutStatus coderws.StatusCode,
timeoutReason string,
) (coderws.MessageType, []byte, error) {
return readOpenAIWSClientMessageWithTimeoutStart(
controlCtx,
conn,
timeout,
timeoutStatus,
timeoutReason,
nil,
nil,
)
}
// readOpenAIWSClientMessageWithTimeoutStart supports readers whose timeout
// starts after a state transition, such as a completed passthrough turn. When
// timeoutActive is nil, a positive timeout starts immediately.
func readOpenAIWSClientMessageWithTimeoutStart(
controlCtx context.Context,
conn *coderws.Conn,
timeout time.Duration,
timeoutStatus coderws.StatusCode,
timeoutReason string,
timeoutStart <-chan struct{},
timeoutActive func() bool,
) (coderws.MessageType, []byte, error) {
if conn == nil {
return 0, nil, errors.New("openai websocket client connection is nil")
}
if controlCtx == nil {
controlCtx = context.Background()
}
readDone := make(chan openAIWSClientReadResult, 1)
go func() {
messageType, payload, err := conn.Read(context.Background())
readDone <- openAIWSClientReadResult{messageType: messageType, payload: payload, err: err}
}()
var timer *time.Timer
var timeoutCh <-chan time.Time
startTimeout := func() {
if timeout <= 0 || (timeoutActive != nil && !timeoutActive()) {
return
}
if timer == nil {
timer = time.NewTimer(timeout)
} else {
if !timer.Stop() {
select {
case <-timer.C:
default:
}
}
timer.Reset(timeout)
}
timeoutCh = timer.C
}
if timeoutActive == nil || timeoutActive() {
startTimeout()
}
defer func() {
if timer != nil {
timer.Stop()
}
}()
closeAndJoin := func(status coderws.StatusCode, reason string, cause error) (coderws.MessageType, []byte, error) {
_ = conn.Close(status, reason)
_ = conn.CloseNow()
<-readDone
return 0, nil, NewOpenAIWSClientCloseError(status, reason, cause)
}
for {
select {
case result := <-readDone:
return result.messageType, result.payload, result.err
case <-timeoutStart:
startTimeout()
case <-timeoutCh:
return closeAndJoin(timeoutStatus, timeoutReason, context.DeadlineExceeded)
case <-controlCtx.Done():
cause := context.Cause(controlCtx)
if errors.Is(cause, ErrOpenAIWSIngressLeaseLost) {
return closeAndJoin(
coderws.StatusTryAgainLater,
"websocket ingress capacity lease lost; please reconnect",
cause,
)
}
return closeAndJoin(coderws.StatusGoingAway, "websocket request canceled", cause)
}
}
}