package proto_socket_test import ( "context" "fmt" "sync" "testing" "time" "nhooyr.io/websocket" toki "git.toki-labs.com/toki/proto-socket/go" "git.toki-labs.com/toki/proto-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 TestWsConcurrentRequests(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() const count = 5 var wg sync.WaitGroup errCh := make(chan error, count) for i := 0; i < count; i++ { i := i wg.Add(1) go func() { defer wg.Done() index := int32(30 + i) message := fmt.Sprintf("request %d", i) res, err := toki.SendRequestTyped[*packets.TestData, *packets.TestData]( &client.Communicator, &packets.TestData{Index: index, Message: message}, 2*time.Second, ) if err != nil { errCh <- err return } if res.GetIndex() != index*2 || res.GetMessage() != "echo: "+message { errCh <- fmt.Errorf("request %d got index=%d message=%q", i, res.GetIndex(), res.GetMessage()) } }() } wg.Wait() close(errCh) for err := range errCh { if err != nil { t.Fatal(err) } } } 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 proto-socket ws"}); err != nil { t.Fatal(err) } select { case msg := <-received: if msg.GetIndex() != 42 || msg.GetMessage() != "hello proto-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") } } func TestWsOriginPatterns(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() // 1. Server with Allowed OriginPatterns portAllowed := freePort(t) optsAllowed := toki.WsServerOptions{ AcceptOptions: &websocket.AcceptOptions{ OriginPatterns: []string{"localhost:3000"}, }, } serverAllowed := toki.NewWsServerWithOptions("127.0.0.1", portAllowed, "/", optsAllowed, func(conn *websocket.Conn) *toki.WsClient { return toki.NewWsClient(conn, 0, 0, testParserMap()) }) if err := serverAllowed.Start(ctx); err != nil { t.Fatal(err) } defer serverAllowed.Stop() // Attempt dial with allowed origin dialOpts := &websocket.DialOptions{ HTTPHeader: map[string][]string{ "Origin": {"http://localhost:3000"}, }, } connAllowed, _, err := websocket.Dial(ctx, fmt.Sprintf("ws://127.0.0.1:%d/", portAllowed), dialOpts) if err != nil { t.Fatalf("expected allowed origin to succeed, got: %v", err) } _ = connAllowed.Close(websocket.StatusNormalClosure, "") // Attempt dial with disallowed origin (should fail) dialOptsBad := &websocket.DialOptions{ HTTPHeader: map[string][]string{ "Origin": {"http://malicious.com"}, }, } _, _, err = websocket.Dial(ctx, fmt.Sprintf("ws://127.0.0.1:%d/", portAllowed), dialOptsBad) if err == nil { t.Fatal("expected disallowed origin to be rejected, but it succeeded") } // 2. Default Server (no options, same-origin only) portDefault := freePort(t) serverDefault := toki.NewWsServer("127.0.0.1", portDefault, "/", func(conn *websocket.Conn) *toki.WsClient { return toki.NewWsClient(conn, 0, 0, testParserMap()) }) if err := serverDefault.Start(ctx); err != nil { t.Fatal(err) } defer serverDefault.Stop() // Default server should reject local cross-origin _, _, err = websocket.Dial(ctx, fmt.Sprintf("ws://127.0.0.1:%d/", portDefault), dialOpts) if err == nil { t.Fatal("expected default server to reject cross-origin, but it succeeded") } } // TestWsReadLoopGatewayParseErrorDisconnects: a malformed binary frame received // by a gateway-attached WsClient.readLoop must disconnect with // DisconnectReasonWSPacketParse, matching the inline parse-error semantics. func TestWsReadLoopGatewayParseErrorDisconnects(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() var mu sync.Mutex var srvClient *toki.WsClient server := toki.NewWsServer("127.0.0.1", port, "/", func(conn *websocket.Conn) *toki.WsClient { c := toki.NewWsClient(conn, 0, 0, testParserMap()) c.EnableInboundGateway(2, 16) mu.Lock() srvClient = c mu.Unlock() return c }) if err := server.Start(ctx); err != nil { t.Fatal(err) } defer server.Stop() // Raw client conn so we can send bytes that are not valid PacketBase. conn, _, err := websocket.Dial(ctx, fmt.Sprintf("ws://127.0.0.1:%d/", port), nil) if err != nil { t.Fatal(err) } defer conn.Close(websocket.StatusNormalClosure, "") if err := conn.Write(ctx, websocket.MessageBinary, []byte{0xff, 0xff, 0xff}); err != nil { t.Fatal(err) } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { mu.Lock() c := srvClient mu.Unlock() if c != nil && c.DisconnectInfo().Reason != toki.DisconnectReasonUnknown { break } time.Sleep(time.Millisecond) } mu.Lock() c := srvClient mu.Unlock() if c == nil { t.Fatal("server never accepted the client connection") } if info := c.DisconnectInfo(); info.Reason != toki.DisconnectReasonWSPacketParse { t.Fatalf("expected %s disconnect via gateway error path, got %q", toki.DisconnectReasonWSPacketParse, info.Reason) } }