- 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
368 lines
9.4 KiB
Go
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)
|
|
}
|
|
}
|