package main import ( "context" "crypto/tls" "crypto/x509" "flag" "fmt" "os" "sync" "time" "google.golang.org/protobuf/proto" toki "git.toki-labs.com/toki/proto-socket/go" "git.toki-labs.com/toki/proto-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 dart server" { fail("2", fmt.Sprintf("unexpected push index=%d message=%q", msg.GetIndex(), msg.GetMessage())) return false } pass("2", "received push from dart 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) }