114 lines
2.2 KiB
Go
Raw Permalink Normal View History

package proxy
import (
"io"
"log"
"net/http"
"net/url"
"strings"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
func newUpgrader(allowedOrigins []string) websocket.Upgrader {
originSet := make(map[string]bool, len(allowedOrigins))
for _, o := range allowedOrigins {
originSet[strings.TrimSpace(o)] = true
}
return websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
if originSet["*"] {
return true
}
origin := r.Header.Get("Origin")
if origin == "" {
return false
}
return originSet[origin]
},
}
}
func NewWSProxy(targetURL string, allowedOrigins []string) gin.HandlerFunc {
upgrader := newUpgrader(allowedOrigins)
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
<-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
}