diff --git a/cookbook/websocket/gorilla/server.go b/cookbook/websocket/gorilla/server.go index 845d3e26..48ec2fd7 100644 --- a/cookbook/websocket/gorilla/server.go +++ b/cookbook/websocket/gorilla/server.go @@ -25,12 +25,14 @@ func hello(c *echo.Context) error { err := ws.WriteMessage(websocket.TextMessage, []byte("Hello, Client!")) if err != nil { c.Logger().Error("failed to write WS message", "error", err) + return nil } // Read _, msg, err := ws.ReadMessage() if err != nil { c.Logger().Error("failed to read WS message", "error", err) + return nil } fmt.Printf("%s\n", msg) } diff --git a/cookbook/websocket/gorilla/server_test.go b/cookbook/websocket/gorilla/server_test.go new file mode 100644 index 00000000..e627ac7b --- /dev/null +++ b/cookbook/websocket/gorilla/server_test.go @@ -0,0 +1,41 @@ +package main + +import ( + "log/slog" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/labstack/echo/v5" +) + +func TestHelloReturnsAfterClientDisconnect(t *testing.T) { + e := echo.New() + e.Logger = slog.New(slog.DiscardHandler) + handlerReturned := make(chan struct{}) + e.GET("/ws", func(c *echo.Context) error { + err := hello(c) + close(handlerReturned) + return err + }) + + server := httptest.NewServer(e) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/ws" + ws, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("failed to connect websocket client: %v", err) + } + if err := ws.Close(); err != nil { + t.Fatalf("failed to close websocket client: %v", err) + } + + select { + case <-handlerReturned: + case <-time.After(2 * time.Second): + t.Fatal("websocket handler did not return after client disconnected") + } +} diff --git a/cookbook/websocket/net/server.go b/cookbook/websocket/net/server.go index 02141114..cbda78a3 100644 --- a/cookbook/websocket/net/server.go +++ b/cookbook/websocket/net/server.go @@ -16,12 +16,14 @@ func hello(c *echo.Context) error { // Write if err := websocket.Message.Send(ws, "Hello, Client!"); err != nil { c.Logger().Error("failed to write WS message", "error", err) + return } // Read msg := "" if err := websocket.Message.Receive(ws, &msg); err != nil { c.Logger().Error("failed to write WS message", "error", err) + return } fmt.Printf("%s\n", msg) } diff --git a/cookbook/websocket/net/server_test.go b/cookbook/websocket/net/server_test.go new file mode 100644 index 00000000..bc82da37 --- /dev/null +++ b/cookbook/websocket/net/server_test.go @@ -0,0 +1,42 @@ +package main + +import ( + "log/slog" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/labstack/echo/v5" + "golang.org/x/net/websocket" +) + +func TestHelloReturnsAfterClientDisconnect(t *testing.T) { + e := echo.New() + e.Logger = slog.New(slog.DiscardHandler) + handlerReturned := make(chan struct{}) + e.GET("/ws", func(c *echo.Context) error { + err := hello(c) + close(handlerReturned) + return err + }) + + server := httptest.NewServer(e) + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/ws" + ws, err := websocket.Dial(wsURL, "", server.URL) + if err != nil { + server.Close() + t.Fatalf("failed to connect websocket client: %v", err) + } + if err := ws.Close(); err != nil { + server.CloseClientConnections() + t.Fatalf("failed to close websocket client: %v", err) + } + + select { + case <-handlerReturned: + server.Close() + case <-time.After(2 * time.Second): + t.Fatal("websocket handler did not return after client disconnected") + } +}