package toki_socket import ( "context" "crypto/ecdsa" "crypto/elliptic" "crypto/rand" "crypto/tls" "crypto/x509" "crypto/x509/pkix" "encoding/binary" "encoding/pem" "io" "math/big" "net" "testing" "time" "google.golang.org/protobuf/proto" "nhooyr.io/websocket" "toki-labs.com/toki_socket/go/packets" ) func TestTcpRequestResponse(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := NewTcpServer("127.0.0.1", port, func(conn net.Conn) *TcpClient { return NewTcpClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *TcpClient) { 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 := DialTcp(ctx, "127.0.0.1", port, 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client.Close() res, err := 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 := NewTcpServer("127.0.0.1", port, func(conn net.Conn) *TcpClient { return NewTcpClient(conn, 0, 0, testParserMap()) }) if err := server.Start(ctx); err != nil { t.Fatal(err) } defer server.Stop() client, err := 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) AddListenerTyped[*packets.TestData](&client.Communicator, func(m *packets.TestData) { received <- m }) for deadline := time.Now().Add(time.Second); time.Now().Before(deadline) && len(server.Clients()) == 0; { time.Sleep(time.Millisecond) } 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 TestTcpBroadcastAttemptsAllClients(t *testing.T) { badConn, badPeer := net.Pipe() badClient := NewTcpClient(badConn, 0, 0, testParserMap()) _ = badClient.Close() _ = badPeer.Close() goodConn, goodPeer := net.Pipe() defer goodPeer.Close() goodClient := NewTcpClient(goodConn, 0, 0, testParserMap()) defer goodClient.Close() received := make(chan *packets.PacketBase, 1) go func() { base, err := readTCPPacket(goodPeer) if err == nil { received <- base } }() server := &TcpServer{clients: []*TcpClient{badClient, goodClient}} err := server.Broadcast(&packets.TestData{Index: 9, Message: "best effort"}) if err == nil { t.Fatal("expected broadcast to report the failed client") } select { case base := <-received: if base.GetTypeName() != TypeNameOf(&packets.TestData{}) { t.Fatalf("unexpected packet type: %s", base.GetTypeName()) } case <-time.After(2 * time.Second): t.Fatal("broadcast did not attempt the healthy client") } } func TestWsRequestResponse(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := NewWsServer("127.0.0.1", port, "/", func(conn *websocket.Conn) *WsClient { return NewWsClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *WsClient) { 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 := DialWsWithHeartbeat(ctx, "127.0.0.1", port, "/", 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client.Close() res, err := 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 TestTcpSendReceive(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() received := make(chan *packets.TestData, 1) server := NewTcpServer("127.0.0.1", port, func(conn net.Conn) *TcpClient { return NewTcpClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *TcpClient) { 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 := 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 toki-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") } } func TestHeartbeatDisconnectsWithoutResponse(t *testing.T) { clientConn, peerConn := net.Pipe() defer peerConn.Close() go func() { _, _ = io.Copy(io.Discard, peerConn) }() client := NewTcpClient(clientConn, 1, 1, testParserMap()) defer client.Close() disconnected := make(chan struct{}, 1) client.AddDisconnectListener(func(*TcpClient) { disconnected <- struct{}{} }) select { case <-disconnected: case <-time.After(3500 * time.Millisecond): t.Fatal("heartbeat timeout did not disconnect client") } if client.IsAlive() { t.Fatal("client is still marked alive after heartbeat timeout") } } func TestTypeNameMatchesDartConvention(t *testing.T) { if got := TypeNameOf(&packets.TestData{}); got != "TestData" { t.Fatalf("type name = %q, want TestData", got) } } func freePort(t *testing.T) int { t.Helper() ln, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatal(err) } defer ln.Close() return ln.Addr().(*net.TCPAddr).Port } func readTCPPacket(conn net.Conn) (*packets.PacketBase, error) { header := make([]byte, 4) if _, err := io.ReadFull(conn, header); err != nil { return nil, err } length := binary.BigEndian.Uint32(header) packetBytes := make([]byte, int(length)) if _, err := io.ReadFull(conn, packetBytes); err != nil { return nil, err } base := &packets.PacketBase{} if err := proto.Unmarshal(packetBytes, base); err != nil { return nil, err } return base, nil } // generateSelfSignedCert creates an in-memory self-signed certificate valid // for 127.0.0.1. Returns the server tls.Config and a client tls.Config that // trusts the generated certificate. func generateSelfSignedCert(t *testing.T) (serverCfg, clientCfg *tls.Config) { t.Helper() priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) if err != nil { t.Fatal(err) } tmpl := &x509.Certificate{ SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "toki-socket-test"}, NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, } certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &priv.PublicKey, priv) if err != nil { t.Fatal(err) } privDER, err := x509.MarshalECPrivateKey(priv) if err != nil { t.Fatal(err) } certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: privDER}) cert, err := tls.X509KeyPair(certPEM, keyPEM) if err != nil { t.Fatal(err) } serverCfg = &tls.Config{Certificates: []tls.Certificate{cert}} pool := x509.NewCertPool() parsed, err := x509.ParseCertificate(certDER) if err != nil { t.Fatal(err) } pool.AddCert(parsed) clientCfg = &tls.Config{RootCAs: pool, ServerName: "127.0.0.1"} return serverCfg, clientCfg } func TestTLSTcp(t *testing.T) { port := freePort(t) ctx, cancel := context.WithCancel(context.Background()) defer cancel() serverTLS, clientTLS := generateSelfSignedCert(t) server := NewTcpServerTLS("127.0.0.1", port, serverTLS, func(conn net.Conn) *TcpClient { return NewTcpClient(conn, 0, 0, testParserMap()) }) server.OnClientConnected = func(client *TcpClient) { AddRequestListenerTyped[*packets.TestData, *packets.TestData](&client.Communicator, func(req *packets.TestData) (*packets.TestData, error) { return &packets.TestData{ Index: req.GetIndex() * 2, Message: "tls-echo: " + req.GetMessage(), }, nil }) } if err := server.Start(ctx); err != nil { t.Fatal(err) } defer server.Stop() client, err := DialTcpTLS(ctx, "127.0.0.1", port, clientTLS, 0, 0, testParserMap()) if err != nil { t.Fatal(err) } defer client.Close() res, err := SendRequestTyped[*packets.TestData, *packets.TestData]( &client.Communicator, &packets.TestData{Index: 7, Message: "hello tls"}, 2*time.Second, ) if err != nil { t.Fatal(err) } if res.GetIndex() != 14 || res.GetMessage() != "tls-echo: hello tls" { t.Fatalf("unexpected response: %v", res) } }