proto-socket/go/test/ws_test.go
toki 2ec43f2d00 feat: inbound queue ordering - Go gateway implementation
- Add inbound_gateway.go and inbound_gateway_test.go
- Update communicator.go with nonce handling improvements
- Update tcp_client.go and ws_client.go
- Update communicator_test.go and ws_test.go
- Move old Go gateway docs to archive
2026-06-02 11:43:10 +09:00

361 lines
10 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")
}
}
// 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)
}
}