- Add browser WebSocket import compatibility test - Add platform-specific WebSocket clients (io/web) - Add protobuf client/server web variants - Fix WS protobuf server for browser compatibility - Remove deprecated ws_protobuf_client.dart (consolidated into platform-specific files) - Update test and server files for websocket functionality
305 lines
8.4 KiB
Go
305 lines
8.4 KiB
Go
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")
|
|
}
|
|
}
|