package proxy import ( "io" "log" "net/http" "net/url" "strings" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" ) var upgrader = websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, } func NewWSProxy(targetURL string) gin.HandlerFunc { return func(c *gin.Context) { clientConn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { log.Printf("[ws_proxy] upgrade error: %v", err) return } defer clientConn.Close() backendURL, err := buildWSURL(targetURL, c.Request) if err != nil { log.Printf("[ws_proxy] url parse error: %v", err) return } header := http.Header{} for _, key := range []string{"X-User-ID", "X-Company-ID", "X-User-Role"} { if v := c.Request.Header.Get(key); v != "" { header.Set(key, v) } } backendConn, _, err := websocket.DefaultDialer.Dial(backendURL, header) if err != nil { log.Printf("[ws_proxy] dial backend error: %v", err) return } defer backendConn.Close() errCh := make(chan error, 2) go pumpMessages(clientConn, backendConn, errCh) go pumpMessages(backendConn, clientConn, errCh) <-errCh } } func pumpMessages(src, dst *websocket.Conn, errCh chan<- error) { for { msgType, reader, err := src.NextReader() if err != nil { errCh <- err return } writer, err := dst.NextWriter(msgType) if err != nil { errCh <- err return } if _, err := io.Copy(writer, reader); err != nil { errCh <- err return } if err := writer.Close(); err != nil { errCh <- err return } } } func buildWSURL(target string, req *http.Request) (string, error) { u, err := url.Parse(target) if err != nil { return "", err } scheme := "ws" if strings.HasPrefix(u.Scheme, "https") { scheme = "wss" } u.Scheme = scheme u.Path = req.URL.Path u.RawQuery = req.URL.RawQuery return u.String(), nil }