proto-socket/go/socket_test.go
toki 14e4cd56d0 feat: add WebSocket/WSS support and Go implementation
- Add WebSocket binary frame support to protocol specification
- Update README and PROTOCOL.md to reflect dual TCP/WebSocket transport
- Add Go implementation with TCP, WebSocket, and heartbeat support
- Include .claude settings configuration
2026-04-11 08:33:46 +09:00

368 lines
9.4 KiB
Go

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)
}
}