* fix: avoid stale stream writes after client disconnect * fix: wait for stream ping goroutines before returning * fix: log stream results after goroutine cleanup * fix: broadcast stream stop signals * fix: abort upstream on client disconnect and restore write error contracts Keep the goroutine-lifecycle fix (unconditional wg.Wait before returning the gin.Context, close resp.Body inside cleanup), but drop the drain-on-disconnect behavior: when the client goes away, cleanup now runs immediately so the upstream body is closed, the provider stops generating, and users are not billed for tokens produced after they disconnected. Also restore FlushWriter/StringData/PingData returning an error when the request context is done, so non-scanner relay loops (ollama, fake-stream, audio, image) keep their disconnect awareness instead of silently consuming the upstream to completion. ResponseChunkData now propagates write errors. Add a bounded per-write deadline (http.NewResponseController) before each locked stream write so a slow-but-connected client cannot block a write forever and hang the unconditional wg.Wait. --------- Co-authored-by: CaIon <i@caion.me>
228 lines
6.1 KiB
Go
228 lines
6.1 KiB
Go
package helper
|
|
|
|
import (
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
"github.com/QuantumNous/new-api/dto"
|
|
"github.com/QuantumNous/new-api/logger"
|
|
"github.com/QuantumNous/new-api/types"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/gorilla/websocket"
|
|
)
|
|
|
|
func FlushWriter(c *gin.Context) (err error) {
|
|
defer func() {
|
|
if r := recover(); r != nil {
|
|
err = fmt.Errorf("flush panic recovered: %v", r)
|
|
}
|
|
}()
|
|
|
|
if c == nil || c.Writer == nil {
|
|
return nil
|
|
}
|
|
|
|
if requestContextDone(c) {
|
|
return fmt.Errorf("request context done: %w", c.Request.Context().Err())
|
|
}
|
|
|
|
flusher, ok := c.Writer.(http.Flusher)
|
|
if !ok {
|
|
return errors.New("streaming error: flusher not found")
|
|
}
|
|
|
|
flusher.Flush()
|
|
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 {
|
|
return
|
|
}
|
|
|
|
// 设置标志,表示头部已经设置过
|
|
c.Set("event_stream_headers_set", true)
|
|
|
|
c.Writer.Header().Set("Content-Type", "text/event-stream")
|
|
c.Writer.Header().Set("Cache-Control", "no-cache")
|
|
c.Writer.Header().Set("Connection", "keep-alive")
|
|
c.Writer.Header().Set("Transfer-Encoding", "chunked")
|
|
c.Writer.Header().Set("X-Accel-Buffering", "no")
|
|
}
|
|
|
|
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())
|
|
} else {
|
|
c.Render(-1, common.CustomEvent{Data: fmt.Sprintf("event: %s\n", resp.Type)})
|
|
c.Render(-1, common.CustomEvent{Data: "data: " + string(jsonData)})
|
|
}
|
|
_ = FlushWriter(c)
|
|
return nil
|
|
}
|
|
|
|
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) 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)})
|
|
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 requestContextDone(c) {
|
|
return fmt.Errorf("request context done: %w", c.Request.Context().Err())
|
|
}
|
|
|
|
c.Render(-1, common.CustomEvent{Data: "data: " + str})
|
|
return FlushWriter(c)
|
|
}
|
|
|
|
func PingData(c *gin.Context) error {
|
|
if c == nil || c.Writer == nil {
|
|
return errors.New("context or writer is nil")
|
|
}
|
|
|
|
if requestContextDone(c) {
|
|
return fmt.Errorf("request context done: %w", c.Request.Context().Err())
|
|
}
|
|
|
|
if _, err := c.Writer.Write([]byte(": PING\n\n")); err != nil {
|
|
return fmt.Errorf("write ping data failed: %w", err)
|
|
}
|
|
return FlushWriter(c)
|
|
}
|
|
|
|
func ObjectData(c *gin.Context, object interface{}) error {
|
|
if object == nil {
|
|
return errors.New("object is nil")
|
|
}
|
|
jsonData, err := common.Marshal(object)
|
|
if err != nil {
|
|
return fmt.Errorf("error marshalling object: %w", err)
|
|
}
|
|
return StringData(c, string(jsonData))
|
|
}
|
|
|
|
func Done(c *gin.Context) {
|
|
_ = StringData(c, "[DONE]")
|
|
}
|
|
|
|
func WssString(c *gin.Context, ws *websocket.Conn, str string) error {
|
|
if ws == nil {
|
|
logger.LogError(c, "websocket connection is nil")
|
|
return errors.New("websocket connection is nil")
|
|
}
|
|
//common.LogInfo(c, fmt.Sprintf("sending message: %s", str))
|
|
return ws.WriteMessage(1, []byte(str))
|
|
}
|
|
|
|
func WssObject(c *gin.Context, ws *websocket.Conn, object interface{}) error {
|
|
jsonData, err := common.Marshal(object)
|
|
if err != nil {
|
|
return fmt.Errorf("error marshalling object: %w", err)
|
|
}
|
|
if ws == nil {
|
|
logger.LogError(c, "websocket connection is nil")
|
|
return errors.New("websocket connection is nil")
|
|
}
|
|
//common.LogInfo(c, fmt.Sprintf("sending message: %s", jsonData))
|
|
return ws.WriteMessage(1, jsonData)
|
|
}
|
|
|
|
func WssError(c *gin.Context, ws *websocket.Conn, openaiError types.OpenAIError) {
|
|
if ws == nil {
|
|
return
|
|
}
|
|
errorObj := &dto.RealtimeEvent{
|
|
Type: "error",
|
|
EventId: GetLocalRealtimeID(c),
|
|
Error: &openaiError,
|
|
}
|
|
_ = WssObject(c, ws, errorObj)
|
|
}
|
|
|
|
func GetResponseID(c *gin.Context) string {
|
|
logID := c.GetString(common.RequestIdKey)
|
|
return fmt.Sprintf("chatcmpl-%s", logID)
|
|
}
|
|
|
|
func GetLocalRealtimeID(c *gin.Context) string {
|
|
logID := c.GetString(common.RequestIdKey)
|
|
return fmt.Sprintf("evt_%s", logID)
|
|
}
|
|
|
|
func GenerateStartEmptyResponse(id string, createAt int64, model string, systemFingerprint *string) *dto.ChatCompletionsStreamResponse {
|
|
return &dto.ChatCompletionsStreamResponse{
|
|
Id: id,
|
|
Object: "chat.completion.chunk",
|
|
Created: createAt,
|
|
Model: model,
|
|
SystemFingerprint: systemFingerprint,
|
|
Choices: []dto.ChatCompletionsStreamResponseChoice{
|
|
{
|
|
Delta: dto.ChatCompletionsStreamResponseChoiceDelta{
|
|
Role: "assistant",
|
|
Content: common.GetPointer(""),
|
|
},
|
|
},
|
|
},
|
|
}
|
|
}
|
|
|
|
func GenerateStopResponse(id string, createAt int64, model string, finishReason string) *dto.ChatCompletionsStreamResponse {
|
|
return &dto.ChatCompletionsStreamResponse{
|
|
Id: id,
|
|
Object: "chat.completion.chunk",
|
|
Created: createAt,
|
|
Model: model,
|
|
SystemFingerprint: nil,
|
|
Choices: []dto.ChatCompletionsStreamResponseChoice{
|
|
{
|
|
FinishReason: &finishReason,
|
|
},
|
|
},
|
|
}
|
|
}
|
|
|
|
func GenerateFinalUsageResponse(id string, createAt int64, model string, usage dto.Usage) *dto.ChatCompletionsStreamResponse {
|
|
return &dto.ChatCompletionsStreamResponse{
|
|
Id: id,
|
|
Object: "chat.completion.chunk",
|
|
Created: createAt,
|
|
Model: model,
|
|
SystemFingerprint: nil,
|
|
Choices: make([]dto.ChatCompletionsStreamResponseChoice, 0),
|
|
Usage: &usage,
|
|
}
|
|
}
|