package proto_socket import ( "sync" "testing" "time" "google.golang.org/protobuf/proto" "git.toki-labs.com/toki/proto-socket/go/packets" ) type nonceTestTransport struct { mu sync.Mutex packets []*packets.PacketBase } func (t *nonceTestTransport) WritePacket(base *packets.PacketBase) error { t.mu.Lock() defer t.mu.Unlock() t.packets = append(t.packets, base) return nil } func (t *nonceTestTransport) Close() error { return nil } func (t *nonceTestTransport) sent() []*packets.PacketBase { t.mu.Lock() defer t.mu.Unlock() return append([]*packets.PacketBase{}, t.packets...) } func nonceTestParserMap() ParserMap { return ParserMap{ TypeNameOf(&packets.TestData{}): func(b []byte) (proto.Message, error) { m := &packets.TestData{} return m, proto.Unmarshal(b, m) }, } } func TestNonceWrapsAfterInt32MaxWithoutEmittingZero(t *testing.T) { transport := &nonceTestTransport{} communicator := NewCommunicator(transport, nonceTestParserMap()) defer communicator.Close() communicator.nonce.Store(MaxNonce - 1) if err := communicator.Send(&packets.TestData{Index: 1}); err != nil { t.Fatal(err) } if err := communicator.Send(&packets.TestData{Index: 2}); err != nil { t.Fatal(err) } sent := transport.sent() if len(sent) != 2 { t.Fatalf("expected 2 packets, got %d", len(sent)) } if got := sent[0].GetNonce(); got != MaxNonce { t.Fatalf("first nonce = %d, want %d", got, MaxNonce) } if got := sent[1].GetNonce(); got != 1 { t.Fatalf("second nonce = %d, want 1", got) } for _, packet := range sent { if packet.GetNonce() == 0 { t.Fatal("emitted reserved nonce 0") } } } func TestNextNonceSkipsActivePendingAfterWrap(t *testing.T) { transport := &nonceTestTransport{} communicator := NewCommunicator(transport, nonceTestParserMap()) defer communicator.Close() // Set nonce to MaxNonce so the next call wraps to 1 communicator.nonce.Store(MaxNonce) // Start a request whose nonce will be 1 (after wrap); leave it unanswered reqErrCh := make(chan error, 1) go func() { _, err := SendRequestTyped[*packets.TestData, *packets.TestData]( communicator, &packets.TestData{Index: 1, Message: "pending"}, time.Second, ) reqErrCh <- err }() // Wait for the request packet with nonce 1 to be sent var req1 *packets.PacketBase for deadline := time.Now().Add(time.Second); time.Now().Before(deadline); { sent := transport.sent() if len(sent) >= 1 { req1 = sent[0] break } time.Sleep(time.Millisecond) } if req1 == nil { t.Fatal("first request packet not sent") } if req1.GetNonce() != 1 { t.Fatalf("first nonce = %d, want 1", req1.GetNonce()) } // Reset nonce to MaxNonce to simulate another wrap cycle communicator.nonce.Store(MaxNonce) // Send must skip nonce 1 (still pending) and use nonce 2 if err := communicator.Send(&packets.TestData{Index: 2}); err != nil { t.Fatal(err) } sent := transport.sent() if len(sent) < 2 { t.Fatalf("expected at least 2 packets, got %d", len(sent)) } if got := sent[len(sent)-1].GetNonce(); got != 2 { t.Fatalf("second nonce = %d, want 2", got) } // Resolve the pending request so the goroutine exits cleanly response := &packets.TestData{Index: 99, Message: "done"} data, _ := proto.Marshal(response) communicator.OnReceivedData(TypeNameOf(response), data, 0, req1.GetNonce()) select { case err := <-reqErrCh: if err != nil { t.Fatalf("pending request error: %v", err) } case <-time.After(time.Second): t.Fatal("pending request did not complete") } } func TestSendRequestNonceWrapsAtInt32Max(t *testing.T) { transport := &nonceTestTransport{} communicator := NewCommunicator(transport, nonceTestParserMap()) defer communicator.Close() communicator.nonce.Store(MaxNonce - 1) resultCh := make(chan *packets.TestData, 1) errCh := make(chan error, 1) go func() { res, err := SendRequestTyped[*packets.TestData, *packets.TestData]( communicator, &packets.TestData{Index: 1, Message: "max"}, time.Second, ) if err != nil { errCh <- err return } resultCh <- res }() var request *packets.PacketBase for deadline := time.Now().Add(time.Second); time.Now().Before(deadline); { sent := transport.sent() if len(sent) == 1 { request = sent[0] break } time.Sleep(time.Millisecond) } if request == nil { t.Fatal("request packet was not sent") } if request.GetNonce() != MaxNonce { t.Fatalf("request nonce = %d, want %d", request.GetNonce(), MaxNonce) } if request.GetNonce() == 0 { t.Fatal("emitted reserved nonce 0") } response := &packets.TestData{Index: 2, Message: "max response"} data, err := proto.Marshal(response) if err != nil { t.Fatal(err) } communicator.OnReceivedData(TypeNameOf(response), data, 0, request.GetNonce()) select { case err := <-errCh: t.Fatal(err) case got := <-resultCh: if got.GetMessage() != "max response" { t.Fatalf("response message = %q, want %q", got.GetMessage(), "max response") } case <-time.After(time.Second): t.Fatal("timed out waiting for response") } }