You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Go语言反向代理中如何正确修改WebSocket消息体

原实现存在的核心问题

  • WebSocket是基于帧的协议,不是裸字节流,直接在Read/Write方法里做字节替换会破坏WebSocket帧结构,且单次读取可能拿到不完整的消息分片,导致替换失效、连接断开。
  • Read/Write方法返回值逻辑错误:无论实际读写多少字节都返回传入缓冲区长度,会导致IO层读取到错误长度的数据,出现丢包、粘包;Write方法逻辑完全写反,先创建空缓冲区写入后端连接,根本没有把客户端发送的真实数据传过去。
  • 类型断言失败时仅打印日志,未做错误处理,后续调用originBody的方法会直接触发空指针panic。
  • Go标准库httputil.ReverseProxy处理WebSocket升级请求时,会通过Hijacker接口劫持TCP连接做双向转发,直接包装resp.Body根本拿不到升级后的双向数据流,包装逻辑不会生效。

正确实现方案

核心思路:接管WebSocket升级后的双向连接,使用成熟的WebSocket库解析完整帧,拿到完整消息负载后做内容替换,替换完成后重新封装为WebSocket帧转发给对端,不要直接操作裸字节流。

前置依赖

先安装WebSocket工具库:

go get github.com/gorilla/websocket

核心实现代码

package rever

import (
	"bytes"
	"context"
	"crypto/tls"
	"github.com/gorilla/websocket"
	"io"
	"log"
	"net/http"
	"net/http/httputil"
	"net/url"
	"strconv"
)

var websocketUpgrader = websocket.Upgrader{
	CheckOrigin: func(r *http.Request) bool {
		// 允许所有跨域请求,生产环境请根据实际安全规则调整
		return true
	},
}

var websocketDialer = &websocket.Dialer{
	TLSClientConfig: &tls.Config{
		InsecureSkipVerify: true,
	},
}

// 消息内容替换逻辑
func rewriteWSMessage(msg []byte, direction string) []byte {
	if direction == "backend2client" {
		// 服务端发往客户端:替换域名为本地域名
		return bytes.ReplaceAll(msg, []byte("mm.remote"), []byte("mm.local"))
	}
	// 客户端发往服务端:替换域名为远端域名
	return bytes.ReplaceAll(msg, []byte("mm.local"), []byte("mm.remote"))
}

// 双向转发WebSocket流量
func proxyWSConn(clientConn *websocket.Conn, backendConn *websocket.Conn) {
	defer func() {
		_ = clientConn.Close()
		_ = backendConn.Close()
	}()

	// 转发后端到客户端的消息
	go func() {
		for {
			msgType, msg, err := backendConn.ReadMessage()
			if err != nil {
				if !websocket.IsCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
					log.Printf("read backend ws message error: %v", err)
				}
				return
			}
			// ping/pong/close等控制帧直接转发,不修改
			if msgType != websocket.TextMessage && msgType != websocket.BinaryMessage {
				_ = clientConn.WriteMessage(msgType, msg)
				continue
			}
			newMsg := rewriteWSMessage(msg, "backend2client")
			if err = clientConn.WriteMessage(msgType, newMsg); err != nil {
				log.Printf("write to client ws error: %v", err)
				return
			}
		}
	}()

	// 转发客户端到后端的消息
	for {
		msgType, msg, err := clientConn.ReadMessage()
		if err != nil {
			if !websocket.IsCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
				log.Printf("read client ws message error: %v", err)
			}
			return
		}
		if msgType != websocket.TextMessage && msgType != websocket.BinaryMessage {
			_ = backendConn.WriteMessage(msgType, msg)
			continue
		}
		newMsg := rewriteWSMessage(msg, "client2backend")
		if err = backendConn.WriteMessage(msgType, newMsg); err != nil {
			log.Printf("write to backend ws error: %v", err)
			return
		}
	}
}

type ctxKey struct{}
func getResponseWriterFromReq(req *http.Request) http.ResponseWriter {
	return req.Context().Value(ctxKey{}).(http.ResponseWriter)
}

type transport struct {
	http.RoundTripper
	backendURL *url.URL
}

func (t *transport) RoundTrip(req *http.Request) (resp *http.Response, err error) {
	// 识别WebSocket升级请求
	isWSReq := websocket.IsWebSocketUpgrade(req)
	if isWSReq {
		// 升级和客户端的WebSocket连接
		clientConn, err := websocketUpgrader.Upgrade(getResponseWriterFromReq(req), req, nil)
		if err != nil {
			log.Printf("upgrade client ws error: %v", err)
			return nil, err
		}
		// 拼接后端WebSocket地址
		backendWSURL := *t.backendURL
		switch backendWSURL.Scheme {
		case "https":
			backendWSURL.Scheme = "wss"
		case "http":
			backendWSURL.Scheme = "ws"
		}
		backendWSURL.Path = req.URL.Path
		backendWSURL.RawQuery = req.URL.RawQuery
		// 建立和后端的WebSocket连接
		backendConn, _, err := websocketDialer.Dial(backendWSURL.String(), req.Header)
		if err != nil {
			log.Printf("dial backend ws error: %v", err)
			_ = clientConn.Close()
			return nil, err
		}
		// 启动双向转发
		go proxyWSConn(clientConn, backendConn)
		// 返回101响应结束HTTP处理流程
		resp = &http.Response{
			StatusCode: http.StatusSwitchingProtocols,
			Body:       io.NopCloser(bytes.NewBuffer(nil)),
			Header:     http.Header{},
		}
		return resp, nil
	}

	// 普通HTTP请求走原有修改逻辑
	resp, err = t.RoundTripper.RoundTrip(req)
	if err != nil {
		log.Println("round trip error: ", err.Error())
		return nil, err
	}

	b, err := io.ReadAll(resp.Body)
	if err != nil {
		log.Println("read response body error: ", err.Error())
		return nil, err
	}
	_ = resp.Body.Close()

	b = bytes.ReplaceAll(b, []byte("mm.remote"), []byte("mm.local"))

	body := io.NopCloser(bytes.NewReader(b))
	resp.Body = body
	resp.ContentLength = int64(len(b))
	resp.Header.Set("Content-Length", strconv.Itoa(len(b)))

	return resp, nil
}

var _ http.RoundTripper = &transport{}

func NewProxy(targetHost string) (*httputil.ReverseProxy, error) {
	backendURL, err := url.Parse(targetHost)
	if err != nil {
		log.Println("parse target url error: ", err.Error())
		return nil, err
	}

	proxy := httputil.NewSingleHostReverseProxy(backendURL)

	originalDirector := proxy.Director
	proxy.Director = func(req *http.Request) {
		originalDirector(req)
		modifyRequest(req)
	}

	proxy.ErrorHandler = func(w http.ResponseWriter, req *http.Request, err error) {
		log.Printf("proxy error: %v", err)
	}

	dt := http.DefaultTransport.(*http.Transport).Clone()
	dt.TLSClientConfig = &tls.Config{InsecureSkipVerify: true}
	dt.ForceAttemptHTTP2 = false
	proxy.Transport = &transport{
		RoundTripper: dt,
		backendURL:   backendURL,
	}

	// 包装ServeHTTP,将ResponseWriter存入请求上下文供WebSocket升级使用
	originServeHTTP := proxy.ServeHTTP
	proxy.ServeHTTP = func(w http.ResponseWriter, r *http.Request) {
		ctx := context.WithValue(r.Context(), ctxKey{}, w)
		originServeHTTP(w, r.WithContext(ctx))
	}

	return proxy, nil
}

func modifyRequest(req *http.Request) {
	req.Host = "mm.remote"
	req.Header.Set("Accept-Encoding", "identity")
}

func ProxyRequestHandler(proxy *httputil.ReverseProxy) func(http.ResponseWriter, *http.Request) {
	return func(w http.ResponseWriter, r *http.Request) {
		proxy.ServeHTTP(w, r)
	}
}

func Main() {
	proxy, err := NewProxy("https://mm.remote")
	if err != nil {
		log.Println("init proxy error: ", err.Error())
		panic(err)
	}

	http.HandleFunc("/", ProxyRequestHandler(proxy))
	log.Println("Server started on :8008")
	log.Fatal(http.ListenAndServe(":8008", nil))
}

注意事项

  • 上述实现默认自动拼接WebSocket分片消息,不需要手动处理分片逻辑,控制帧直接转发避免连接异常。
  • 原代码使用的ioutil包在Go 1.16版本后已废弃,相关函数已迁移至io包,直接调用io包对应方法即可。
  • 生产环境请调整跨域校验规则,不要直接放开所有来源,避免出现安全风险。
  • 如果需要处理二进制消息的内容修改,直接在rewriteWSMessage方法中补充对应逻辑即可。

内容的提问来源于stack exchange,提问作者Anji

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.28 19:27:23