Files
token_thief/proxy/proxy.go
T

379 lines
11 KiB
Go

package proxy
import (
"context"
"crypto/tls"
"errors"
"log"
"net"
"net/http"
"net/http/httputil"
"net/netip"
"net/url"
"strings"
"sync"
"sync/atomic"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/logger"
)
// LogSubmitter 是 proxy 唯一依赖的日志接收方接口。
type LogSubmitter interface {
Submit(*logger.LogEntry)
}
// Handler 构造反代 HTTP handler。
type Handler struct {
rp *httputil.ReverseProxy
filter *config.Filter
queue LogSubmitter
maxBodyBytes int64
maxRequestBytes int64
trusted []netip.Prefix
connMu sync.Mutex
conns map[net.Conn]struct{}
connChanged chan struct{}
closing bool
closed bool
}
// Options 控制反代连接上游时的网络行为。
type Options struct {
UpstreamTimeout time.Duration
ResponseTimeout time.Duration
SSEIdleTimeout time.Duration
UpstreamTLSInsecureSkipVerify bool
TrustedProxies []netip.Prefix
MaxRequestBytes int64
}
// requestState 通过 context 在 ErrorHandler / ModifyResponse / 主 handler 之间共享状态。
type requestState struct {
requestID string
lastErr atomic.Pointer[string]
upgrade atomic.Pointer[upgradeResponse]
readFailed atomic.Bool
}
type upgradeResponse struct {
status int
header http.Header
}
type ctxKey struct{}
func newRequestState(id string) *requestState { return &requestState{requestID: id} }
func stateFromCtx(ctx context.Context) *requestState {
v, _ := ctx.Value(ctxKey{}).(*requestState)
return v
}
func New(upstream *url.URL, filter *config.Filter, queue LogSubmitter, maxBody int64, upstreamTimeout ...time.Duration) *Handler {
opts := Options{}
if len(upstreamTimeout) > 0 {
opts.UpstreamTimeout = upstreamTimeout[0]
}
return NewWithOptions(upstream, filter, queue, maxBody, opts)
}
func NewWithOptions(upstream *url.URL, filter *config.Filter, queue LogSubmitter, maxBody int64, opts Options) *Handler {
if opts.MaxRequestBytes <= 0 {
opts.MaxRequestBytes = 16 << 20
}
rp := httputil.NewSingleHostReverseProxy(upstream)
rp.FlushInterval = -1 // 让流式 chunk 立即转发
if opts.UpstreamTimeout > 0 || opts.UpstreamTLSInsecureSkipVerify {
transport := http.DefaultTransport.(*http.Transport).Clone()
if opts.UpstreamTimeout > 0 {
dialer := &net.Dialer{Timeout: opts.UpstreamTimeout, KeepAlive: 30 * time.Second}
transport.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) {
conn, err := dialer.DialContext(ctx, network, address)
if err != nil {
return nil, err
}
return &writeTimeoutConn{Conn: conn, timeout: opts.UpstreamTimeout}, nil
}
transport.ResponseHeaderTimeout = opts.UpstreamTimeout
transport.TLSHandshakeTimeout = opts.UpstreamTimeout
}
if opts.UpstreamTLSInsecureSkipVerify {
transport.TLSClientConfig = &tls.Config{InsecureSkipVerify: true}
}
rp.Transport = transport
}
origDirector := rp.Director
rp.Director = func(r *http.Request) {
origDirector(r)
r.Host = upstream.Host
}
// ModifyResponse 在响应头写回客户端之前调用,确保 X-Request-Id 一定生效。
rp.ModifyResponse = func(resp *http.Response) error {
if st := stateFromCtx(resp.Request.Context()); st != nil {
resp.Header.Set("X-Request-Id", st.requestID)
if resp.StatusCode == http.StatusSwitchingProtocols &&
strings.EqualFold(resp.Request.Header.Get("Upgrade"), "websocket") &&
strings.EqualFold(resp.Header.Get("Upgrade"), "websocket") {
st.upgrade.Store(&upgradeResponse{status: resp.StatusCode, header: resp.Header.Clone()})
}
}
if timeout := responseBodyTimeout(resp, opts); timeout > 0 {
resp.Body = newTimeoutBody(resp.Body, timeout, isStreamResponse(resp.Header))
}
if st := stateFromCtx(resp.Request.Context()); st != nil && resp.StatusCode != http.StatusSwitchingProtocols {
resp.Body = &trackingBody{ReadCloser: resp.Body, failed: &st.readFailed}
}
return nil
}
rp.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) {
log.Printf("[proxy] upstream error for %s %s: %v", r.Method, r.URL.Path, err)
// 把错误暴露给主 handler,使其能写入日志。
if st := stateFromCtx(r.Context()); st != nil {
s := err.Error()
st.lastErr.Store(&s)
// ErrorHandler 路径下 ModifyResponse 不会被调用,这里手动写 X-Request-Id。
w.Header().Set("X-Request-Id", st.requestID)
}
http.Error(w, "bad gateway", http.StatusBadGateway)
}
return &Handler{
rp: rp,
filter: filter,
queue: queue,
maxBodyBytes: maxBody,
maxRequestBytes: opts.MaxRequestBytes,
trusted: append([]netip.Prefix(nil), opts.TrustedProxies...),
conns: make(map[net.Conn]struct{}),
connChanged: make(chan struct{}),
}
}
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// 健康检查不参与反代与日志。
if r.URL.Path == "/healthz" {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("ok"))
return
}
shouldLog := h.filter.ShouldLog(r.URL.Path)
reqBody, reqTruncated, err := readRequestBody(r, h.maxBodyBytes, h.maxRequestBytes)
if err != nil && !shouldLog {
log.Printf("[proxy] read request body failed: %v", err)
_ = r.Body.Close()
if errors.Is(err, errRequestBodyTooLarge) {
http.Error(w, "request body too large", http.StatusRequestEntityTooLarge)
} else {
http.Error(w, "bad request", http.StatusBadRequest)
}
return
}
if !shouldLog {
log.Printf("[proxy] skip log by filter method=%s path=%s", r.Method, r.URL.Path)
cw := newCaptureWriter(w, 0)
cw.OnHijack(h.trackConn)
h.rp.ServeHTTP(cw, r)
return
}
started := time.Now()
reqID := newRequestID()
log.Printf("[proxy] capture start request_id=%s method=%s path=%s", reqID, r.Method, r.URL.Path)
st := newRequestState(reqID)
r = r.WithContext(context.WithValue(r.Context(), ctxKey{}, st))
cw := newCaptureWriter(w, h.maxBodyBytes)
if err != nil {
log.Printf("[proxy] read request body failed: %v", err)
_ = r.Body.Close()
s := "read request body: " + err.Error()
st.lastErr.Store(&s)
h.serveRequestBodyError(cw, r, st, started, reqID, errors.Is(err, errRequestBodyTooLarge))
return
}
reqHeadersJSON := headersJSON(r.Header)
clientAddr := clientIP(r, h.trusted)
method := r.Method
path := r.URL.Path
query := r.URL.RawQuery
var submitOnce sync.Once
submit := func(finished time.Time) {
entry, ok := h.buildLogEntry(cw, st, logEntryInput{
requestID: reqID,
method: method,
path: path,
query: query,
clientAddr: clientAddr,
requestHeaders: reqHeadersJSON,
requestBody: reqBody,
requestTruncated: reqTruncated,
started: started,
finished: finished,
})
if !ok {
return
}
h.queue.Submit(entry)
log.Printf("[proxy] capture finish request_id=%s method=%s path=%s status=%d is_stream=%v latency_ms=%d",
reqID, method, path, entry.StatusCode, entry.IsStream, entry.LatencyMS)
}
cw.OnHijack(func(conn net.Conn) net.Conn {
tracked := h.trackConn(conn)
if upgrade := st.upgrade.Load(); upgrade != nil {
cw.SetHijackedResponse(upgrade.status, upgrade.header)
return &headerDeliveryConn{
Conn: tracked,
onDelivered: func() {
submitOnce.Do(func() { submit(time.Now()) })
},
}
}
return tracked
})
h.rp.ServeHTTP(cw, r)
finished := time.Now()
cw.ConfirmDelivery()
responseComplete := !isStreamResponse(cw.Header()) || cw.SSEComplete()
if cw.Complete() && responseComplete && !st.readFailed.Load() && r.Context().Err() == nil {
submitOnce.Do(func() { submit(finished) })
}
}
func (h *Handler) serveRequestBodyError(cw *captureWriter, r *http.Request, st *requestState, started time.Time, requestID string, tooLarge bool) {
cw.Header().Set("X-Request-Id", requestID)
status := http.StatusBadRequest
message := "bad request"
if tooLarge {
status = http.StatusRequestEntityTooLarge
message = "request body too large"
}
http.Error(cw, message, status)
if !cw.Complete() {
return
}
entry, ok := h.buildLogEntry(cw, st, logEntryInput{
requestID: requestID,
method: r.Method,
path: r.URL.Path,
query: r.URL.RawQuery,
clientAddr: clientIP(r, h.trusted),
requestHeaders: headersJSON(r.Header),
started: started,
finished: time.Now(),
})
if ok {
h.queue.Submit(entry)
}
}
// Shutdown closes all active hijacked connections and waits for their release.
func (h *Handler) Shutdown(ctx context.Context) error {
for {
h.connMu.Lock()
h.closing = true
if len(h.conns) == 0 {
h.closed = true
h.connMu.Unlock()
return nil
}
conns := make([]net.Conn, 0, len(h.conns))
for conn := range h.conns {
conns = append(conns, conn)
}
changed := h.connChanged
h.connMu.Unlock()
for _, conn := range conns {
_ = conn.Close()
}
select {
case <-changed:
case <-ctx.Done():
return ctx.Err()
}
}
}
func (h *Handler) trackConn(conn net.Conn) net.Conn {
tracked := &trackedConn{Conn: conn}
tracked.onClose = func() {
h.connMu.Lock()
delete(h.conns, tracked)
close(h.connChanged)
h.connChanged = make(chan struct{})
h.connMu.Unlock()
}
h.connMu.Lock()
if h.closing || h.closed {
h.connMu.Unlock()
_ = conn.Close()
return conn
}
h.conns[tracked] = struct{}{}
close(h.connChanged)
h.connChanged = make(chan struct{})
h.connMu.Unlock()
return tracked
}
type logEntryInput struct {
requestID string
method string
path string
query string
clientAddr string
requestHeaders []byte
requestBody []byte
requestTruncated bool
started time.Time
finished time.Time
}
func (h *Handler) buildLogEntry(cw *captureWriter, st *requestState, in logEntryInput) (*logger.LogEntry, bool) {
if cw.Hijacked() && cw.Status() != http.StatusSwitchingProtocols {
return nil, false
}
var errMsg string
if p := st.lastErr.Load(); p != nil {
errMsg = *p
}
isStream := isStreamResponse(cw.Header())
responseBody := append([]byte(nil), cw.Body()...)
responseTruncated := cw.Truncated()
if isStream && !responseTruncated {
if assembled, ok := assembleSSEJSON(responseBody); ok {
responseBody = assembled
}
}
return &logger.LogEntry{
RequestID: in.requestID,
Method: in.method,
Path: in.path,
Query: in.query,
ClientIP: in.clientAddr,
RequestHeaders: in.requestHeaders,
RequestBody: in.requestBody,
RequestTruncated: in.requestTruncated,
StatusCode: cw.Status(),
ResponseHeaders: headersJSON(cw.Header()),
ResponseBody: responseBody,
ResponseTruncated: responseTruncated,
IsStream: isStream,
LatencyMS: in.finished.Sub(in.started).Milliseconds(),
StartedAt: in.started,
FinishedAt: in.finished,
Error: errMsg,
}, true
}