package proto_socket_test import ( "context" "fmt" "net" "sync" "testing" "time" toki "git.toki-labs.com/toki/proto-socket/go" "git.toki-labs.com/toki/proto-socket/go/packets" ) func TestTcpRequestResponse(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := toki.NewTcpServer("127.0.0.1", port, func(conn net.Conn) *toki.TcpClient { return toki.NewTcpClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *toki.TcpClient) { 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.DialTcp(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"}, 2*time.Second, ) if err != nil { t.Fatal(err) } if res.GetIndex() != 42 || res.GetMessage() != "echo: hello" { t.Fatalf("unexpected response: %v", res) } } func TestTcpBroadcast(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := toki.NewTcpServer("127.0.0.1", port, func(conn net.Conn) *toki.TcpClient { return toki.NewTcpClient(conn, 0, 0, testParserMap()) }) if err := server.Start(ctx); err != nil { t.Fatal(err) } client, err := toki.DialTcp(ctx, "127.0.0.1", port, 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client.Close() received := make(chan *packets.TestData, 1) toki.AddListenerTyped[*packets.TestData](&client.Communicator, func(m *packets.TestData) { received <- m }) waitForCondition(t, time.Second, func() bool { return len(server.Clients()) == 1 }, "client did not connect") if err := server.Broadcast(&packets.TestData{Index: 9, Message: "broadcast"}); err != nil { t.Fatal(err) } select { case msg := <-received: if msg.GetMessage() != "broadcast" { t.Fatalf("unexpected broadcast: %v", msg) } case <-time.After(2 * time.Second): t.Fatal("broadcast timed out") } } func TestTcpServerStopDisconnectsClients(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := toki.NewTcpServer("127.0.0.1", port, func(conn net.Conn) *toki.TcpClient { return toki.NewTcpClient(conn, 0, 0, testParserMap()) }) if err := server.Start(ctx); err != nil { t.Fatal(err) } client, err := toki.DialTcp(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.TcpClient) { 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 TestTcpClientCloseIdempotent(t *testing.T) { clientConn, peerConn := net.Pipe() defer peerConn.Close() client := toki.NewTcpClient(clientConn, 0, 0, testParserMap()) for i := 0; i < 3; i++ { if err := client.Close(); err != nil && i == 0 { t.Fatal(err) } } if client.IsAlive() { t.Fatal("client is still marked alive after close") } } func TestTcpConcurrentRequests(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := toki.NewTcpServer("127.0.0.1", port, func(conn net.Conn) *toki.TcpClient { return toki.NewTcpClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *toki.TcpClient) { 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.DialTcp(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 TestTcpSendReceive(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() received := make(chan *packets.TestData, 1) server := toki.NewTcpServer("127.0.0.1", port, func(conn net.Conn) *toki.TcpClient { return toki.NewTcpClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *toki.TcpClient) { toki.AddListenerTyped[*packets.TestData](&client.Communicator, func(m *packets.TestData) { received <- m }) } if err := server.Start(ctx); err != nil { t.Fatal(err) } defer server.Stop() client, err := toki.DialTcp(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"}); err != nil { t.Fatal(err) } select { case msg := <-received: if msg.GetIndex() != 42 { t.Fatalf("unexpected message: %v", msg) } case <-time.After(2 * time.Second): t.Fatal("receive timed out") } }