proto-socket/go/crosstest/kotlin_go_client/main.go

205 lines
4.8 KiB
Go

package main
import (
"context"
"flag"
"fmt"
"os"
"sync"
"time"
"google.golang.org/protobuf/proto"
toki "toki-labs.com/toki_socket/go"
"toki-labs.com/toki_socket/go/packets"
)
const (
host = "127.0.0.1"
wsPath = "/"
connectWindow = 3 * time.Second
requestWindow = 2 * time.Second
)
type clientHandle struct {
communicator *toki.Communicator
send func(proto.Message) error
close func() error
}
func parserMap() toki.ParserMap {
return toki.ParserMap{
toki.TypeNameOf(&packets.TestData{}): func(b []byte) (proto.Message, error) {
m := &packets.TestData{}
return m, proto.Unmarshal(b, m)
},
}
}
func main() {
mode := flag.String("mode", "tcp", "transport mode: tcp or ws")
port := flag.Int("port", 0, "server port")
phase := flag.String("phase", "send-push", "test phase: send-push or requests")
flag.Parse()
fmt.Printf("INFO typeName go=%s\n", toki.TypeNameOf(&packets.TestData{}))
if *port == 0 {
fail("setup", "port is required")
os.Exit(1)
}
client, err := dialWithRetry(*mode, *port)
if err != nil {
fail("setup", err.Error())
os.Exit(1)
}
defer client.close()
var ok bool
switch *phase {
case "send-push":
ok = runSendPush(client)
case "requests":
ok = runRequests(client)
default:
fail("setup", fmt.Sprintf("unknown phase %q", *phase))
ok = false
}
if !ok {
os.Exit(1)
}
}
func dialWithRetry(mode string, port int) (*clientHandle, error) {
deadline := time.Now().Add(connectWindow)
var lastErr error
for time.Now().Before(deadline) {
ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond)
handle, err := dial(ctx, mode, port)
cancel()
if err == nil {
return handle, nil
}
lastErr = err
time.Sleep(100 * time.Millisecond)
}
return nil, fmt.Errorf("connect %s:%d timed out: %w", mode, port, lastErr)
}
func dial(ctx context.Context, mode string, port int) (*clientHandle, error) {
switch mode {
case "tcp":
client, err := toki.DialTcp(ctx, host, port, 0, 0, parserMap())
if err != nil {
return nil, err
}
return &clientHandle{
communicator: &client.Communicator,
send: client.Send,
close: client.Close,
}, nil
case "ws":
client, err := toki.DialWsWithHeartbeat(ctx, host, port, wsPath, 0, 0, parserMap())
if err != nil {
return nil, err
}
return &clientHandle{
communicator: &client.Communicator,
send: client.Send,
close: client.Close,
}, nil
default:
return nil, fmt.Errorf("unknown mode %q", mode)
}
}
func runSendPush(client *clientHandle) bool {
pushCh := make(chan *packets.TestData, 1)
toki.AddListenerTyped[*packets.TestData](client.communicator, func(msg *packets.TestData) {
pushCh <- msg
})
err := client.send(&packets.TestData{
Index: 101,
Message: "fire from go client",
})
if err != nil {
fail("1", err.Error())
return false
}
pass("1", "fire-and-forget sent")
select {
case msg := <-pushCh:
if msg.GetIndex() != 200 || msg.GetMessage() != "push from kotlin server" {
fail("2", fmt.Sprintf("unexpected push index=%d message=%q", msg.GetIndex(), msg.GetMessage()))
return false
}
pass("2", "received push from kotlin server")
return true
case <-time.After(requestWindow):
fail("2", "timeout waiting for server push")
return false
}
}
func runRequests(client *clientHandle) bool {
res, err := toki.SendRequestTyped[*packets.TestData, *packets.TestData](
client.communicator,
&packets.TestData{Index: 21, Message: "single request from go"},
requestWindow,
)
if err != nil {
fail("3", err.Error())
return false
}
if res.GetIndex() != 42 || res.GetMessage() != "echo: single request from go" {
fail("3", fmt.Sprintf("unexpected response index=%d message=%q", res.GetIndex(), res.GetMessage()))
return false
}
pass("3", "single request response matched")
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("multi request %d from go", i)
res, err := toki.SendRequestTyped[*packets.TestData, *packets.TestData](
client.communicator,
&packets.TestData{Index: index, Message: message},
requestWindow,
)
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 {
fail("4", err.Error())
return false
}
}
pass("4", "concurrent request responses matched")
return true
}
func pass(scenario, detail string) {
fmt.Printf("PASS scenario=%s detail=%s\n", scenario, detail)
}
func fail(scenario, detail string) {
fmt.Printf("FAIL scenario=%s error=%s\n", scenario, detail)
}