package proto_socket_test import ( "strings" "sync" "testing" "time" "google.golang.org/protobuf/proto" toki "git.toki-labs.com/toki/common-proto-socket/go" "git.toki-labs.com/toki/common-proto-socket/go/packets" ) type fakeTransport struct { mu sync.Mutex packets []*packets.PacketBase } func (f *fakeTransport) WritePacket(base *packets.PacketBase) error { f.mu.Lock() defer f.mu.Unlock() f.packets = append(f.packets, base) return nil } func (f *fakeTransport) Close() error { return nil } func (f *fakeTransport) sent() []*packets.PacketBase { f.mu.Lock() defer f.mu.Unlock() return append([]*packets.PacketBase{}, f.packets...) } func testParserMap() 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 TestSendRequestTypeMismatch(t *testing.T) { transport := &fakeTransport{} communicator := toki.NewCommunicator(transport, testParserMap()) defer communicator.Close() errCh := make(chan error, 1) go func() { _, err := toki.SendRequestTyped[*packets.TestData, *packets.TestData]( communicator, &packets.TestData{Index: 1, Message: "hello"}, time.Second, ) errCh <- err }() var requestNonce int32 for deadline := time.Now().Add(time.Second); time.Now().Before(deadline); { sent := transport.sent() if len(sent) == 1 { requestNonce = sent[0].GetNonce() break } time.Sleep(time.Millisecond) } if requestNonce == 0 { t.Fatal("request packet was not sent") } hb, err := proto.Marshal(&packets.HeartBeat{}) if err != nil { t.Fatal(err) } communicator.OnReceivedData(toki.TypeNameOf(&packets.HeartBeat{}), hb, 0, requestNonce) err = <-errCh if err == nil { t.Fatal("expected response type mismatch error") } } func TestSendRequestTimeout(t *testing.T) { transport := &fakeTransport{} communicator := toki.NewCommunicator(transport, testParserMap()) defer communicator.Close() _, err := toki.SendRequestTyped[*packets.TestData, *packets.TestData]( communicator, &packets.TestData{Index: 1, Message: "no response"}, 25*time.Millisecond, ) if err == nil { t.Fatal("expected request timeout error") } if !strings.Contains(err.Error(), "timeout") { t.Fatalf("expected timeout error, got %v", err) } if len(transport.sent()) != 1 { t.Fatalf("expected one request packet to be sent, got %d", len(transport.sent())) } } func TestListenerAndRequestListenerConflict(t *testing.T) { transport := &fakeTransport{} communicator := toki.NewCommunicator(transport, testParserMap()) defer communicator.Close() toki.AddListenerTyped[*packets.TestData](communicator, func(*packets.TestData) {}) defer func() { if recover() == nil { t.Fatal("expected panic") } }() toki.AddRequestListenerTyped[*packets.TestData, *packets.TestData](communicator, func(req *packets.TestData) (*packets.TestData, error) { return req, nil }) }