Files
token_thief/proxy/capture.go
T

114 lines
2.7 KiB
Go

package proxy
import (
"bytes"
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"io"
"net/http"
"net/netip"
"strings"
)
var errRequestBodyTooLarge = errors.New("request body too large")
// readRequestBody 在转发前完整读取请求体,确保读取失败时不会向上游发送损坏请求。
func readRequestBody(r *http.Request, max, requestLimit int64) (captured []byte, truncated bool, err error) {
if r.Body == nil || r.Body == http.NoBody {
return nil, false, nil
}
if r.ContentLength > requestLimit {
return nil, false, errRequestBodyTooLarge
}
body, err := io.ReadAll(io.LimitReader(r.Body, requestLimit+1))
if err != nil {
return nil, false, err
}
if int64(len(body)) > requestLimit {
return nil, false, errRequestBodyTooLarge
}
if err := r.Body.Close(); err != nil {
return nil, false, err
}
r.Body = io.NopCloser(bytes.NewReader(body))
if int64(len(body)) > max {
return body[:max], true, nil
}
return body, false, nil
}
func headersJSON(h http.Header) []byte {
if len(h) == 0 {
return nil
}
b, err := json.Marshal(h)
if err != nil {
return nil
}
return b
}
func newRequestID() string {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "unknown"
}
return hex.EncodeToString(b[:])
}
// clientIP 从请求中提取客户端 IP。
func clientIP(r *http.Request, trusted []netip.Prefix) string {
peer, ok := parsePeerAddr(r.RemoteAddr)
if !ok {
return r.RemoteAddr
}
if !isTrusted(peer, trusted) {
return peer.String()
}
xff := strings.Split(r.Header.Get("X-Forwarded-For"), ",")
if len(xff) == 1 && strings.TrimSpace(xff[0]) == "" {
if realIP, err := netip.ParseAddr(strings.TrimSpace(r.Header.Get("X-Real-IP"))); err == nil {
return realIP.Unmap().String()
}
return peer.String()
}
chain := make([]netip.Addr, len(xff))
for i, raw := range xff {
addr, err := netip.ParseAddr(strings.TrimSpace(raw))
if err != nil {
return peer.String()
}
chain[i] = addr.Unmap()
}
client := peer
for i := len(chain) - 1; i >= 0 && isTrusted(client, trusted); i-- {
client = chain[i]
}
return client.String()
}
func parsePeerAddr(remote string) (netip.Addr, bool) {
if addrPort, err := netip.ParseAddrPort(remote); err == nil {
return addrPort.Addr().Unmap(), true
}
addr, err := netip.ParseAddr(remote)
return addr.Unmap(), err == nil
}
func isTrusted(addr netip.Addr, prefixes []netip.Prefix) bool {
for _, prefix := range prefixes {
if prefix.Contains(addr) {
return true
}
}
return false
}
// isStreamResponse 通过响应头判断是否为流式响应。
func isStreamResponse(h http.Header) bool {
ct := strings.ToLower(strings.TrimSpace(strings.SplitN(h.Get("Content-Type"), ";", 2)[0]))
return ct == "text/event-stream"
}