diff --git a/core/init/router/router.go b/core/init/router/router.go index 51bb8760c205..b83a644befa5 100644 --- a/core/init/router/router.go +++ b/core/init/router/router.go @@ -87,6 +87,7 @@ func Routers() *gin.Engine { } Router.Use(middleware.FrontendFallback()) + Router.Use(middleware.WebSocketOriginGuard()) Router.Use(middleware.OperationLog()) Router.Use(middleware.GlobalLoading()) Router.Use(xpack.AuthProvider.CoreAPIAuthMiddleware()) diff --git a/core/middleware/websocket_origin.go b/core/middleware/websocket_origin.go new file mode 100644 index 000000000000..3756d1e8496b --- /dev/null +++ b/core/middleware/websocket_origin.go @@ -0,0 +1,41 @@ +package middleware + +import ( + "net/http" + "net/url" + "strings" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +func WebSocketOriginGuard() gin.HandlerFunc { + return func(c *gin.Context) { + if !websocket.IsWebSocketUpgrade(c.Request) { + c.Next() + return + } + + if site := c.GetHeader("Sec-Fetch-Site"); site != "" && site != "same-origin" { + c.AbortWithStatus(http.StatusForbidden) + return + } + origins := c.Request.Header.Values("Origin") + if len(origins) == 0 { + c.Next() + return + } + if len(origins) != 1 { + c.AbortWithStatus(http.StatusForbidden) + return + } + origin, err := url.Parse(origins[0]) + if err != nil || (origin.Scheme != "http" && origin.Scheme != "https") || + origin.Host == "" || !strings.EqualFold(origin.Host, c.Request.Host) || + origin.User != nil || origin.Path != "" || origin.RawQuery != "" || origin.ForceQuery || origin.Fragment != "" { + c.AbortWithStatus(http.StatusForbidden) + return + } + c.Next() + } +} diff --git a/frontend/vite.config.ts b/frontend/vite.config.ts index 1ca708da2878..1761725fd1f5 100644 --- a/frontend/vite.config.ts +++ b/frontend/vite.config.ts @@ -88,6 +88,13 @@ export default defineConfig(async ({ mode }: ConfigEnv): Promise => target: 'http://localhost:9999/', changeOrigin: true, ws: true, + configure(proxy) { + proxy.on('proxyReqWs', (proxyReq, req) => { + if (req.headers.host) { + proxyReq.setHeader('Host', req.headers.host); + } + }); + }, }, }, },