Files
token_thief/proxy/writer.go
T

174 lines
3.8 KiB
Go

package proxy
import (
"bufio"
"bytes"
"errors"
"net"
"net/http"
"sync"
)
// captureWriter 包装 http.ResponseWriter,边转发边缓冲响应体。
// 实现 http.Flusher 与 http.Hijacker 以支持 SSE/chunked/WebSocket。
type captureWriter struct {
http.ResponseWriter
buf bytes.Buffer
max int64
written int64
truncated bool
status int
wroteHeader bool
hijacked bool
writeFailed bool
onHijack func(net.Conn) net.Conn
sse sseEventTracker
}
func newCaptureWriter(w http.ResponseWriter, max int64) *captureWriter {
return &captureWriter{ResponseWriter: w, max: max, status: http.StatusOK, sse: newSSEEventTracker(max)}
}
func (c *captureWriter) WriteHeader(code int) {
if c.wroteHeader {
return
}
c.status = code
c.wroteHeader = true
c.ResponseWriter.WriteHeader(code)
}
func (c *captureWriter) Write(p []byte) (int, error) {
if !c.wroteHeader {
c.wroteHeader = true
}
n, err := c.ResponseWriter.Write(p)
if err != nil || n != len(p) {
c.writeFailed = true
}
if n > len(p) {
n = len(p)
}
if n > 0 {
// 仅缓冲 max 字节以内的内容。
remaining := c.max - c.written
if remaining > 0 {
toBuf := n
if int64(toBuf) > remaining {
toBuf = int(remaining)
c.truncated = true
}
c.buf.Write(p[:toBuf])
} else if c.max > 0 {
c.truncated = true
}
c.written += int64(n)
if c.max > 0 && isStreamResponse(c.Header()) {
c.sse.Write(p[:n])
}
c.flush()
}
return n, err
}
func (c *captureWriter) Flush() {
c.flush()
}
// ConfirmDelivery establishes an observable delivery boundary for responses
// whose headers were not followed by a body write.
func (c *captureWriter) ConfirmDelivery() {
if c.hijacked || c.written != 0 {
return
}
if _, ok := c.ResponseWriter.(http.Flusher); !ok {
c.writeFailed = true
return
}
c.flush()
}
func (c *captureWriter) flush() {
if c.hijacked {
return
}
if _, ok := c.ResponseWriter.(http.Flusher); !ok {
return
}
if err := http.NewResponseController(c.ResponseWriter).Flush(); err != nil {
c.writeFailed = true
}
}
func (c *captureWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
h, ok := c.ResponseWriter.(http.Hijacker)
if !ok {
return nil, nil, errors.New("hijack not supported")
}
conn, rw, err := h.Hijack()
if err != nil {
c.writeFailed = true
return nil, nil, err
}
c.hijacked = true
if c.onHijack != nil {
conn = c.onHijack(conn)
rw = bufio.NewReadWriter(rw.Reader, bufio.NewWriter(conn))
}
return conn, rw, nil
}
func (c *captureWriter) Body() []byte { return c.buf.Bytes() }
func (c *captureWriter) Truncated() bool { return c.truncated }
func (c *captureWriter) Status() int { return c.status }
func (c *captureWriter) Hijacked() bool { return c.hijacked }
func (c *captureWriter) Complete() bool { return !c.writeFailed }
func (c *captureWriter) SSEComplete() bool { return c.sse.Complete() }
func (c *captureWriter) OnHijack(fn func(net.Conn) net.Conn) { c.onHijack = fn }
func (c *captureWriter) SetHijackedResponse(status int, header http.Header) {
c.status = status
for key := range c.Header() {
c.Header().Del(key)
}
for key, values := range header {
c.Header()[key] = append([]string(nil), values...)
}
}
type trackedConn struct {
net.Conn
closeOnce sync.Once
onClose func()
}
func (c *trackedConn) Close() error {
err := c.Conn.Close()
c.closeOnce.Do(c.onClose)
return err
}
type headerDeliveryConn struct {
net.Conn
tail []byte
onDelivered func()
delivered bool
}
func (c *headerDeliveryConn) Write(p []byte) (int, error) {
n, err := c.Conn.Write(p)
if c.delivered || err != nil || n == 0 {
return n, err
}
c.tail = append(c.tail, p[:n]...)
if bytes.Contains(c.tail, []byte("\r\n\r\n")) {
c.delivered = true
c.tail = nil
c.onDelivered()
} else if len(c.tail) > 3 {
c.tail = append(c.tail[:0], c.tail[len(c.tail)-3:]...)
}
return n, err
}