mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-21 14:19:18 +08:00
Merge pull request #4019 from wucm667/fix/compact-keepalive-writer-nil-guard
fix(service): 补齐 compact keepalive writer nil 守卫
This commit is contained in:
@@ -1,6 +1,9 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -190,52 +193,105 @@ type openAICompactKeepaliveWriter struct {
|
||||
// suspend 停拍心跳;幂等。任何响应构造(含 Header 访问——写响应必先操作
|
||||
// 响应头)都视为请求侧接管 ResponseWriter。
|
||||
func (w *openAICompactKeepaliveWriter) suspend() {
|
||||
if w.k == nil {
|
||||
return
|
||||
}
|
||||
w.k.Stop()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Header() http.Header {
|
||||
w.suspend()
|
||||
if w.ResponseWriter == nil {
|
||||
return http.Header{}
|
||||
}
|
||||
return w.ResponseWriter.Header()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Write(data []byte) (int, error) {
|
||||
w.suspend()
|
||||
if w.ResponseWriter == nil {
|
||||
return 0, nil
|
||||
}
|
||||
return w.ResponseWriter.Write(data)
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) WriteString(s string) (int, error) {
|
||||
w.suspend()
|
||||
if w.ResponseWriter == nil {
|
||||
return 0, nil
|
||||
}
|
||||
return w.ResponseWriter.WriteString(s)
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) WriteHeader(code int) {
|
||||
w.suspend()
|
||||
if w.ResponseWriter == nil {
|
||||
return
|
||||
}
|
||||
w.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) WriteHeaderNow() {
|
||||
w.suspend()
|
||||
if w.ResponseWriter == nil {
|
||||
return
|
||||
}
|
||||
w.ResponseWriter.WriteHeaderNow()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Flush() {
|
||||
w.suspend()
|
||||
if w.ResponseWriter == nil {
|
||||
return
|
||||
}
|
||||
w.ResponseWriter.Flush()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||
if w.ResponseWriter == nil {
|
||||
return nil, nil, errors.New("response writer released")
|
||||
}
|
||||
return w.ResponseWriter.Hijack()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) CloseNotify() <-chan bool {
|
||||
if w.ResponseWriter == nil {
|
||||
ch := make(chan bool)
|
||||
close(ch)
|
||||
return ch
|
||||
}
|
||||
return w.ResponseWriter.CloseNotify()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Pusher() http.Pusher {
|
||||
if w.ResponseWriter == nil {
|
||||
return nil
|
||||
}
|
||||
return w.ResponseWriter.Pusher()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Status() int {
|
||||
if w.k == nil || w.ResponseWriter == nil {
|
||||
return 0
|
||||
}
|
||||
w.k.mu.Lock()
|
||||
defer w.k.mu.Unlock()
|
||||
return w.ResponseWriter.Status()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Size() int {
|
||||
if w.k == nil || w.ResponseWriter == nil {
|
||||
return 0
|
||||
}
|
||||
w.k.mu.Lock()
|
||||
defer w.k.mu.Unlock()
|
||||
return w.ResponseWriter.Size()
|
||||
}
|
||||
|
||||
func (w *openAICompactKeepaliveWriter) Written() bool {
|
||||
if w.k == nil || w.ResponseWriter == nil {
|
||||
return false
|
||||
}
|
||||
w.k.mu.Lock()
|
||||
defer w.k.mu.Unlock()
|
||||
return w.ResponseWriter.Written()
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
@@ -141,6 +142,110 @@ func TestOpenAICompactKeepaliveWriter_RequestSideWriteSuspendsBeats(t *testing.T
|
||||
require.Contains(t, rec.Body.String(), `{"error":"local reject"}`)
|
||||
}
|
||||
|
||||
func TestOpenAICompactKeepaliveWriter_NilInnerWriter_NoPanic(t *testing.T) {
|
||||
w := &openAICompactKeepaliveWriter{
|
||||
k: &openAICompactSSEKeepalive{stop: make(chan struct{})},
|
||||
}
|
||||
w.ResponseWriter = nil
|
||||
|
||||
assert.NotPanics(t, func() {
|
||||
assert.Equal(t, 0, w.Status())
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
assert.Equal(t, 0, w.Size())
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
assert.False(t, w.Written())
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
assert.NotNil(t, w.Header())
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
n, err := w.Write([]byte("test"))
|
||||
assert.Equal(t, 0, n)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
n, err := w.WriteString("test")
|
||||
assert.Equal(t, 0, n)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
w.WriteHeaderNow()
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
w.Flush()
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
conn, rw, err := w.Hijack()
|
||||
assert.Nil(t, conn)
|
||||
assert.Nil(t, rw)
|
||||
assert.Error(t, err)
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
ch := w.CloseNotify()
|
||||
assert.NotNil(t, ch)
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
assert.Nil(t, w.Pusher())
|
||||
})
|
||||
}
|
||||
|
||||
func TestOpenAICompactKeepaliveWriter_NilKeepalive_NoPanic(t *testing.T) {
|
||||
c, rec := newCompactBridgeTestContext(t, true)
|
||||
w := &openAICompactKeepaliveWriter{ResponseWriter: c.Writer}
|
||||
|
||||
assert.NotPanics(t, func() {
|
||||
assert.Equal(t, 0, w.Status())
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
assert.Equal(t, 0, w.Size())
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
assert.False(t, w.Written())
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
w.Header().Set("X-Test", "ok")
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
n, err := w.WriteString("ok")
|
||||
assert.Equal(t, 2, n)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
assert.NotPanics(t, func() {
|
||||
w.Flush()
|
||||
})
|
||||
require.Equal(t, "ok", rec.Header().Get("X-Test"))
|
||||
require.Equal(t, "ok", rec.Body.String())
|
||||
}
|
||||
|
||||
func TestOpenAICompactKeepaliveWriter_DelegatesWhenReady(t *testing.T) {
|
||||
c, rec := newCompactBridgeTestContext(t, true)
|
||||
stop := StartOpenAICompactSSEKeepalive(c, time.Hour)
|
||||
defer stop()
|
||||
|
||||
w, ok := c.Writer.(*openAICompactKeepaliveWriter)
|
||||
require.True(t, ok)
|
||||
|
||||
w.Header().Set("X-Test", "ok")
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
n, err := w.WriteString("ready")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, len("ready"), n)
|
||||
|
||||
require.Equal(t, http.StatusAccepted, w.Status())
|
||||
require.Equal(t, len("ready"), w.Size())
|
||||
require.True(t, w.Written())
|
||||
require.Equal(t, "ok", rec.Header().Get("X-Test"))
|
||||
require.Equal(t, "ready", rec.Body.String())
|
||||
}
|
||||
|
||||
// fast policy block 在心跳提交后必须降级为 response.failed 终止事件。
|
||||
func TestWriteOpenAIFastPolicyBlockedResponse_AfterKeepaliveCommit(t *testing.T) {
|
||||
c, rec := newCompactBridgeTestContext(t, true)
|
||||
|
||||
Reference in New Issue
Block a user