proto-socket/go/test/tcp_test.go
toki d0754a353a refactor: rename toki_socket to proto_socket across all languages
- Rename package/module from toki_socket to proto_socket in Dart, Kotlin, Python
- Update crosstest implementations to use renamed packages
- Add new proto_socket skill, deprecate add-toki-socket-crosstest-language skill
- Update domain rules for all languages
- Update documentation (README, PORTING_GUIDE, PROTOCOL, VERSIONING)
2026-05-02 07:19:12 +09:00

247 lines
6.2 KiB
Go

package proto_socket_test
import (
"context"
"fmt"
"net"
"sync"
"testing"
"time"
toki "git.toki-labs.com/toki/common-proto-socket/go"
"git.toki-labs.com/toki/common-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")
}
}