proto-socket/go/communicator_nonce_test.go

197 lines
4.9 KiB
Go

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")
}
}