package websocket import ( "net/http" "net/http/httptest" "strings" "sync/atomic" "testing" "time" "github.com/gorilla/websocket" melody "github.com/olahol/melody" ) func TestServerReceivesAndSendsTextMessages(t *testing.T) { server := New(Config{}) var firstHandlerCalled atomic.Bool server.OnMessage(func(_ *melody.Session, _ []byte) { firstHandlerCalled.Store(true) }) server.OnMessage(func(session *melody.Session, message []byte) { _ = server.Send(session, append([]byte("echo:"), message...)) }) httpServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { _ = server.HandleRequest(w, r) })) defer httpServer.Close() defer server.Close() url := "ws" + strings.TrimPrefix(httpServer.URL, "http") conn, _, err := websocket.DefaultDialer.Dial(url, nil) if err != nil { t.Fatal(err) } defer conn.Close() if err = conn.WriteMessage(websocket.TextMessage, []byte("hello")); err != nil { t.Fatal(err) } _ = conn.SetReadDeadline(time.Now().Add(time.Second)) _, message, err := conn.ReadMessage() if err != nil { t.Fatal(err) } if string(message) != "echo:hello" { t.Fatalf("message = %q", message) } if !firstHandlerCalled.Load() { t.Fatal("first message handler was overwritten") } } func TestServerAllowsAnyOriginByDefault(t *testing.T) { server := New(Config{}) defer server.Close() request := httptest.NewRequest(http.MethodGet, "http://127.0.0.1:8000/ws", nil) request.Host = "127.0.0.1:8000" request.Header.Set("Origin", "http://evil.example") if !server.Melody().Upgrader.CheckOrigin(request) { t.Fatal("origin was rejected when allow_origins is empty") } }