package toki_socket_test import ( "context" "testing" "time" "nhooyr.io/websocket" toki "toki-labs.com/toki_socket/go" "toki-labs.com/toki_socket/go/packets" ) func TestWsRequestResponse(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := toki.NewWsServer("127.0.0.1", port, "/", func(conn *websocket.Conn) *toki.WsClient { return toki.NewWsClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *toki.WsClient) { toki.AddRequestListenerTyped[*packets.TestData, *packets.TestData](&client.Communicator, func(req *packets.TestData) (*packets.TestData, error) { return &packets.TestData{ Index: req.GetIndex() * 2, Message: "echo: " + req.GetMessage(), }, nil }) } if err := server.Start(ctx); err != nil { t.Fatal(err) } defer server.Stop() client, err := toki.DialWsWithHeartbeat(ctx, "127.0.0.1", port, "/", 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client.Close() res, err := toki.SendRequestTyped[*packets.TestData, *packets.TestData]( &client.Communicator, &packets.TestData{Index: 21, Message: "hello ws"}, 2*time.Second, ) if err != nil { t.Fatal(err) } if res.GetIndex() != 42 || res.GetMessage() != "echo: hello ws" { t.Fatalf("unexpected response: %v", res) } } func TestWsSendReceive(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() received := make(chan *packets.TestData, 1) server := toki.NewWsServer("127.0.0.1", port, "/", func(conn *websocket.Conn) *toki.WsClient { return toki.NewWsClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *toki.WsClient) { toki.AddListenerTyped[*packets.TestData](&client.Communicator, func(m *packets.TestData) { received <- m }) } if err := server.Start(ctx); err != nil { t.Fatal(err) } client, err := toki.DialWsWithHeartbeat(ctx, "127.0.0.1", port, "/", 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client.Close() if err := client.Send(&packets.TestData{Index: 42, Message: "hello toki-socket ws"}); err != nil { t.Fatal(err) } select { case msg := <-received: if msg.GetIndex() != 42 || msg.GetMessage() != "hello toki-socket ws" { t.Fatalf("unexpected message: %v", msg) } case <-time.After(2 * time.Second): t.Fatal("receive timed out") } } func TestWsBroadcast(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := toki.NewWsServer("127.0.0.1", port, "/", func(conn *websocket.Conn) *toki.WsClient { return toki.NewWsClient(conn, 0, 0, testParserMap()) }) if err := server.Start(ctx); err != nil { t.Fatal(err) } defer server.Stop() client1, err := toki.DialWsWithHeartbeat(ctx, "127.0.0.1", port, "/", 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client1.Close() client2, err := toki.DialWsWithHeartbeat(ctx, "127.0.0.1", port, "/", 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client2.Close() received1 := make(chan *packets.TestData, 1) received2 := make(chan *packets.TestData, 1) toki.AddListenerTyped[*packets.TestData](&client1.Communicator, func(m *packets.TestData) { received1 <- m }) toki.AddListenerTyped[*packets.TestData](&client2.Communicator, func(m *packets.TestData) { received2 <- m }) waitForCondition(t, time.Second, func() bool { return len(server.Clients()) == 2 }, "clients did not connect") if err := server.Broadcast(&packets.TestData{Index: 9, Message: "ws broadcast"}); err != nil { t.Fatal(err) } for name, ch := range map[string]chan *packets.TestData{ "client1": received1, "client2": received2, } { select { case msg := <-ch: if msg.GetMessage() != "ws broadcast" { t.Fatalf("%s got unexpected broadcast: %v", name, msg) } case <-time.After(2 * time.Second): t.Fatalf("%s broadcast timed out", name) } } } func TestWsServerStopDisconnectsClients(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := toki.NewWsServer("127.0.0.1", port, "/", func(conn *websocket.Conn) *toki.WsClient { return toki.NewWsClient(conn, 0, 0, testParserMap()) }) if err := server.Start(ctx); err != nil { t.Fatal(err) } client, err := toki.DialWsWithHeartbeat(ctx, "127.0.0.1", port, "/", 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client.Close() disconnected := make(chan struct{}, 1) client.AddDisconnectListener(func(*toki.WsClient) { disconnected <- struct{}{} }) waitForCondition(t, time.Second, func() bool { return len(server.Clients()) == 1 }, "client did not connect") if err := server.Stop(); err != nil { t.Fatal(err) } select { case <-disconnected: case <-time.After(2 * time.Second): t.Fatal("client disconnect callback timed out") } if client.IsAlive() { t.Fatal("client is still marked alive after server stop") } }