Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 18 additions & 26 deletions relay/channel/api_request.go
Original file line number Diff line number Diff line change
Expand Up @@ -396,10 +396,12 @@ func DoWssRequest(a Adaptor, c *gin.Context, info *common.RelayInfo, requestBody
return targetConn, nil
}

func startPingKeepAlive(c *gin.Context, pingInterval time.Duration) context.CancelFunc {
func startPingKeepAlive(c *gin.Context, pingInterval time.Duration) (context.CancelFunc, <-chan struct{}) {
pingerCtx, stopPinger := context.WithCancel(context.Background())
done := make(chan struct{})

gopool.Go(func() {
defer close(done)
defer func() {
// 增加panic恢复处理
if r := recover(); r != nil {
Expand Down Expand Up @@ -449,36 +451,24 @@ func startPingKeepAlive(c *gin.Context, pingInterval time.Duration) context.Canc
}
})

return stopPinger
return stopPinger, done
}

func sendPingData(c *gin.Context, mutex *sync.Mutex) error {
// 增加超时控制,防止锁死等待
done := make(chan error, 1)
go func() {
mutex.Lock()
defer mutex.Unlock()
mutex.Lock()
defer mutex.Unlock()

err := helper.PingData(c)
if err != nil {
logger.LogError(c, "SSE ping error: "+err.Error())
done <- err
return
}

logger.LogDebug(c, "SSE ping data sent")
done <- nil
}()

// 设置发送ping数据的超时时间
select {
case err := <-done:
// Bound the write so a slow client cannot block this goroutine forever;
// doRequest's defer waits for the pinger to exit before returning.
helper.ExtendWriteDeadline(c)
err := helper.PingData(c)
if err != nil {
logger.LogError(c, "SSE ping error: "+err.Error())
return err
case <-time.After(10 * time.Second):
return errors.New("SSE ping data send timeout")
case <-c.Request.Context().Done():
return errors.New("request context cancelled during ping")
}

logger.LogDebug(c, "SSE ping data sent")
return nil
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}

func DoRequest(c *gin.Context, req *http.Request, info *common.RelayInfo) (*http.Response, error) {
Expand All @@ -497,17 +487,19 @@ func doRequest(c *gin.Context, req *http.Request, info *common.RelayInfo) (*http
}

var stopPinger context.CancelFunc
var pingerDone <-chan struct{}
if info.IsStream {
helper.SetEventStreamHeaders(c)
// 处理流式请求的 ping 保活
generalSettings := operation_setting.GetGeneralSetting()
if generalSettings.PingIntervalEnabled && !info.DisablePing {
pingInterval := time.Duration(generalSettings.PingIntervalSeconds) * time.Second
stopPinger = startPingKeepAlive(c, pingInterval)
stopPinger, pingerDone = startPingKeepAlive(c, pingInterval)
// 使用defer确保在任何情况下都能停止ping goroutine
defer func() {
if stopPinger != nil {
stopPinger()
<-pingerDone
logger.LogDebug(c, "SSE ping goroutine stopped by defer")
}
}()
Expand Down
2 changes: 1 addition & 1 deletion relay/channel/openai/helper.go
Original file line number Diff line number Diff line change
Expand Up @@ -206,5 +206,5 @@ func sendResponsesStreamData(c *gin.Context, streamResponse dto.ResponsesStreamR
if data == "" {
return
}
helper.ResponseChunkData(c, streamResponse, data)
_ = helper.ResponseChunkData(c, streamResponse, data)
}
20 changes: 6 additions & 14 deletions relay/channel/openai/relay_image.go
Original file line number Diff line number Diff line change
Expand Up @@ -136,10 +136,10 @@ func writeOpenaiImageStreamChunk(c *gin.Context, data []byte) {
}
_ = common.Unmarshal(data, &payload)
if eventName := strings.TrimSpace(payload.Type); eventName != "" {
c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("event: %s\n", eventName)})
_ = helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: eventName}, string(data))
return
}
c.Render(-1, common.CustomEvent{Data: "data: " + string(data)})
_ = helper.FlushWriter(c)
_ = helper.StringData(c, string(data))
}

// isOpenAIImageStreamErrorEvent detects upstream error chunks by JSON content
Expand Down Expand Up @@ -269,19 +269,11 @@ func writeOpenaiImageStreamPayload(c *gin.Context, eventName string, payload any
return err
}
if eventName != "" {
if _, err := fmt.Fprintf(c.Writer, "event: %s\n", eventName); err != nil {
return err
}
}
if _, err := fmt.Fprintf(c.Writer, "data: %s\n\n", data); err != nil {
return err
return helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: eventName}, string(data))
}
return helper.FlushWriter(c)
return helper.StringData(c, string(data))
}

func writeOpenaiImageStreamDone(c *gin.Context) error {
if _, err := fmt.Fprint(c.Writer, "data: [DONE]\n\n"); err != nil {
return err
}
return helper.FlushWriter(c)
return helper.StringData(c, "[DONE]")
}
26 changes: 21 additions & 5 deletions relay/helper/common.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ func FlushWriter(c *gin.Context) (err error) {
return nil
}

if c.Request != nil && c.Request.Context().Err() != nil {
if requestContextDone(c) {
return fmt.Errorf("request context done: %w", c.Request.Context().Err())
}

Expand All @@ -38,6 +38,10 @@ func FlushWriter(c *gin.Context) (err error) {
return nil
}

func requestContextDone(c *gin.Context) bool {
return c != nil && c.Request != nil && c.Request.Context().Err() != nil
}

func SetEventStreamHeaders(c *gin.Context) {
// 检查是否已经设置过头部
if _, exists := c.Get("event_stream_headers_set"); exists {
Expand All @@ -55,6 +59,10 @@ func SetEventStreamHeaders(c *gin.Context) {
}

func ClaudeData(c *gin.Context, resp dto.ClaudeResponse) error {
if requestContextDone(c) {
return nil
}

jsonData, err := common.Marshal(resp)
if err != nil {
common.SysError("error marshalling stream response: " + err.Error())
Expand All @@ -67,23 +75,31 @@ func ClaudeData(c *gin.Context, resp dto.ClaudeResponse) error {
}

func ClaudeChunkData(c *gin.Context, resp dto.ClaudeResponse, data string) {
if requestContextDone(c) {
return
}

c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("event: %s\n", resp.Type)})
c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("data: %s\n", data)})
_ = FlushWriter(c)
}

func ResponseChunkData(c *gin.Context, resp dto.ResponsesStreamResponse, data string) {
func ResponseChunkData(c *gin.Context, resp dto.ResponsesStreamResponse, data string) error {
if requestContextDone(c) {
return fmt.Errorf("request context done: %w", c.Request.Context().Err())
}

c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("event: %s\n", resp.Type)})
c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("data: %s", data)})
_ = FlushWriter(c)
return FlushWriter(c)
}

func StringData(c *gin.Context, str string) error {
if c == nil || c.Writer == nil {
return errors.New("context or writer is nil")
}

if c.Request != nil && c.Request.Context().Err() != nil {
if requestContextDone(c) {
return fmt.Errorf("request context done: %w", c.Request.Context().Err())
}

Expand All @@ -96,7 +112,7 @@ func PingData(c *gin.Context) error {
return errors.New("context or writer is nil")
}

if c.Request != nil && c.Request.Context().Err() != nil {
if requestContextDone(c) {
return fmt.Errorf("request context done: %w", c.Request.Context().Err())
}

Expand Down
Loading