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
相关产品推荐
相关产品推荐

