254 lines
6.2 KiB
Go
254 lines
6.2 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"crypto/tls"
|
|
"crypto/x509"
|
|
"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, ws, tls, or wss")
|
|
port := flag.Int("port", 0, "server port")
|
|
phase := flag.String("phase", "send-push", "test phase: send-push or requests")
|
|
cert := flag.String("cert", "", "path to PEM certificate for TLS verification")
|
|
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, *cert)
|
|
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 buildClientTLS(certFile string) (*tls.Config, error) {
|
|
certPEM, err := os.ReadFile(certFile)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("read cert %s: %w", certFile, err)
|
|
}
|
|
pool := x509.NewCertPool()
|
|
if !pool.AppendCertsFromPEM(certPEM) {
|
|
return nil, fmt.Errorf("no valid PEM certificate found in %s", certFile)
|
|
}
|
|
return &tls.Config{RootCAs: pool}, nil
|
|
}
|
|
|
|
func dialWithRetry(mode string, port int, certFile string) (*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, certFile)
|
|
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, certFile string) (*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
|
|
case "tls":
|
|
if certFile == "" {
|
|
return nil, fmt.Errorf("--cert is required for tls mode")
|
|
}
|
|
tlsCfg, err := buildClientTLS(certFile)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
client, err := toki.DialTcpTLS(ctx, host, port, tlsCfg, 0, 0, parserMap())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &clientHandle{
|
|
communicator: &client.Communicator,
|
|
send: client.Send,
|
|
close: client.Close,
|
|
}, nil
|
|
case "wss":
|
|
if certFile == "" {
|
|
return nil, fmt.Errorf("--cert is required for wss mode")
|
|
}
|
|
tlsCfg, err := buildClientTLS(certFile)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
client, err := toki.DialWssWithHeartbeat(ctx, host, port, wsPath, tlsCfg, 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 python server" {
|
|
fail("2", fmt.Sprintf("unexpected push index=%d message=%q", msg.GetIndex(), msg.GetMessage()))
|
|
return false
|
|
}
|
|
pass("2", "received push from python 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)
|
|
}
|