Files
xk-hy-forward-go/internal/forward/proxy.go
2026-05-28 17:04:29 +08:00

257 lines
8.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package forward
import (
"bytes"
"context"
"crypto/tls"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"log"
"net"
"net/http"
"net/http/httputil"
"net/url"
"strings"
"time"
"unicode/utf8"
)
const (
outboundDialTimeout = 10 * time.Second
outboundResponseHeaderTimeout = 60 * time.Second
)
// dryRunResponse 为 ENABLE_FORWARD=false 时返回给 transit-go 的 JSON 结构(回显将要转发的出站内容)。
type dryRunResponse struct {
DryRun bool `json:"dry_run"`
Kind string `json:"kind"`
TargetURL string `json:"target_url"`
Method string `json:"method"`
Headers map[string][]string `json:"headers"`
Body string `json:"body,omitempty"`
BodyBase64 string `json:"body_base64"`
}
// newReverseProxy 创建指向 target 的反向代理;Director 内会改写 URL 并过滤 Header 白名单。
func newReverseProxy(target, kind string) (*httputil.ReverseProxy, error) {
// 步骤 1:解析上游固定地址
targetURL, err := url.Parse(target)
if err != nil {
return nil, err
}
proxy := httputil.NewSingleHostReverseProxy(targetURL)
// 步骤 2:政务网自签证书,TLS 跳过校验;出站连接/读头超时便于快速失败
proxy.Transport = &http.Transport{
TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, //nolint:gosec // 政务网自签证书
DialContext: (&net.Dialer{Timeout: outboundDialTimeout}).DialContext,
ResponseHeaderTimeout: outboundResponseHeaderTimeout,
}
// 步骤 3:每次转发前改写出站 URL 与 Header,并记录出站快照供 api-log 第 2 块
proxy.Director = func(req *http.Request) {
applyOutboundTarget(req, targetURL, kind)
if col := apiLogCollectorFrom(req.Context()); col != nil {
col.SetOutbound(req.URL.String(), req.Header)
}
}
// 步骤 4:捕获上游响应供 api-log 第 3 块,Body 回填后继续透传
proxy.ModifyResponse = func(resp *http.Response) error {
body, err := io.ReadAll(resp.Body)
if err != nil {
return err
}
resp.Body = io.NopCloser(bytes.NewReader(body))
if col := apiLogCollectorFrom(resp.Request.Context()); col != nil {
col.SetUpstream(resp.StatusCode, resp.Header, body)
}
return nil
}
return proxy, nil
}
// applyOutboundTarget 将入站请求改写为指向政务云的出站形态(固定 URL + Header 白名单,不读 Body)。
func applyOutboundTarget(req *http.Request, targetURL *url.URL, kind string) {
req.URL.Scheme = targetURL.Scheme
req.URL.Host = targetURL.Host
req.URL.Path = targetURL.Path
req.URL.RawPath = targetURL.RawPath
req.URL.RawQuery = targetURL.RawQuery
req.Host = targetURL.Host
filtered := make(http.Header)
copyAllowedHeaders(filtered, req.Header, kind)
req.Header = filtered
}
// buildOutboundFromInbound 按 Director 相同规则组装出站 URL/Header(不读 Body)。
func buildOutboundFromInbound(headers http.Header, target, kind string) (targetURL string, hdr map[string][]string, err error) {
parsed, err := url.Parse(target)
if err != nil {
return "", nil, err
}
filtered := make(http.Header)
copyAllowedHeaders(filtered, headers, kind)
hdr = make(map[string][]string, len(filtered))
for k, vs := range filtered {
cp := make([]string, len(vs))
copy(cp, vs)
hdr[k] = cp
}
return parsed.String(), hdr, nil
}
// recordOutboundSnapshot 将出站快照写入 api-log 采集器(403/dry-run 等未走 Director 时使用)。
func recordOutboundSnapshot(ctx context.Context, headers http.Header, target, kind string) {
col := apiLogCollectorFrom(ctx)
if col == nil {
return
}
targetURL, hdrMap, err := buildOutboundFromInbound(headers, target, kind)
if err != nil {
return
}
filtered := make(http.Header)
for k, vs := range hdrMap {
for _, v := range vs {
filtered.Add(k, v)
}
}
col.SetOutbound(targetURL, filtered)
}
// buildOutboundSnapshot 读取入站 Body,按与真实转发相同规则组装出站 URL/Header/Body,供回显模式使用。
func buildOutboundSnapshot(r *http.Request, target, kind string) (targetURL string, headers map[string][]string, body []byte, err error) {
body, err = io.ReadAll(r.Body)
if err != nil {
return "", nil, nil, err
}
_ = r.Body.Close()
targetURL, headers, err = buildOutboundFromInbound(r.Header, target, kind)
if err != nil {
return "", nil, nil, err
}
return targetURL, headers, body, nil
}
// handleForward 处理监管/文件两条转发路由:校验 → 回显或真实转发 → 记日志。
func handleForward(w http.ResponseWriter, r *http.Request, kind string, proxy *httputil.ReverseProxy, target, allowIPs, forwardSecret string, enableForward bool) {
if col := apiLogCollectorFrom(r.Context()); col != nil {
col.SetMeta(kind, target)
}
// 步骤 1:仅允许 POST
if r.Method != http.MethodPost {
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
ip := clientIP(r)
// 步骤 2:来源 IP 白名单(ALLOW_IPS 非空时生效)
if allowIPs != "" && !ipAllowed(ip, allowIPs) {
recordOutboundSnapshot(r.Context(), r.Header, target, kind)
appLogf("%s | %s | %s | forbidden | not in ALLOW_IPS", time.Now().Format(time.RFC3339), ip, kind)
http.Error(w, "forbidden", http.StatusForbidden)
return
}
// 步骤 3:内网共享密钥(FORWARD_SHARED_SECRET 非空时校验 X-Forward-Token)
if forwardSecret != "" && r.Header.Get("X-Forward-Token") != forwardSecret {
recordOutboundSnapshot(r.Context(), r.Header, target, kind)
appLogf("%s | %s | %s | forbidden | bad forward token", time.Now().Format(time.RFC3339), ip, kind)
http.Error(w, "forbidden", http.StatusForbidden)
return
}
// 步骤 4:未开启真实转发时,回显将要发往政务云的请求头与请求体
if !enableForward {
handleDryRun(w, r, kind, target, ip)
return
}
// 步骤 5:ReverseProxy 转发至政务云,记录耗时与上游 HTTP 状态
start := time.Now()
sw := &statusRecorder{ResponseWriter: w, status: http.StatusOK}
proxy.ServeHTTP(sw, r)
ms := time.Since(start).Milliseconds()
line := fmt.Sprintf("%s | %s | %s | http=%d | %dms | %s", time.Now().Format(time.RFC3339), ip, kind, sw.status, ms, target)
appLogf(line)
log.Printf("forward %s %s http=%d %dms", ip, kind, sw.status, ms)
}
// handleDryRun 不请求政务云,将按转发规则组装后的出站 URL、Header、Body 以 JSON 返回给 transit-go。
func handleDryRun(w http.ResponseWriter, r *http.Request, kind, target, ip string) {
start := time.Now()
// 步骤 1:按与 Director 相同规则组装出站快照
targetURL, headers, body, err := buildOutboundSnapshot(r, target, kind)
if err != nil {
appLogf("%s | %s | %s | dry_run_error | %v", time.Now().Format(time.RFC3339), ip, kind, err)
http.Error(w, "internal error", http.StatusInternalServerError)
return
}
recordOutboundSnapshot(r.Context(), r.Header, target, kind)
// 步骤 2:构造 JSON(body 明文仅 UTF-8 且非 multipart 时填充;body_base64 始终提供)
resp := dryRunResponse{
DryRun: true,
Kind: kind,
TargetURL: targetURL,
Method: http.MethodPost,
Headers: headers,
BodyBase64: base64.StdEncoding.EncodeToString(body),
}
if plainBodyForDryRun(kind, headers, body) {
resp.Body = string(body)
}
// 步骤 3:写回 transit-go
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(http.StatusOK)
if err := json.NewEncoder(w).Encode(resp); err != nil {
appLogf("%s | %s | %s | dry_run_encode_error | %v", time.Now().Format(time.RFC3339), ip, kind, err)
return
}
ms := time.Since(start).Milliseconds()
line := fmt.Sprintf("%s | %s | %s | dry_run | target=%s | body_len=%d | %dms", time.Now().Format(time.RFC3339), ip, kind, targetURL, len(body), ms)
appLogf(line)
log.Printf("dry_run %s %s body_len=%d %dms", ip, kind, len(body), ms)
}
// plainBodyForDryRun 判断是否可在 JSON 的 body 字段中放明文(file/multipart 或非法 UTF-8 则仅 base64)。
func plainBodyForDryRun(kind string, headers map[string][]string, body []byte) bool {
if kind == "file" {
return false
}
ct := ""
if vs, ok := headers["Content-Type"]; ok && len(vs) > 0 {
ct = strings.ToLower(vs[0])
}
if strings.Contains(ct, "multipart/") {
return false
}
return utf8.Valid(body)
}
// statusRecorder 包装 ResponseWriter,记录上游返回的 HTTP 状态码供日志使用。
type statusRecorder struct {
http.ResponseWriter
status int
}
func (w *statusRecorder) WriteHeader(code int) {
w.status = code
w.ResponseWriter.WriteHeader(code)
}
// clientIP 从 RemoteAddr 解析客户端 IP(去掉端口)。
func clientIP(r *http.Request) string {
ip, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return ip
}
// ipAllowed 判断 ip 是否在逗号分隔的 ALLOW_IPS 列表中(精确匹配)。
func ipAllowed(ip, allow string) bool {
for _, part := range strings.Split(allow, ",") {
if strings.TrimSpace(part) == ip {
return true
}
}
return false
}