package notify import ( "bytes" "context" "encoding/binary" "fmt" "io" "net" "sync" "testing" "time" ) type delayedRecordPacket struct { data []byte ready time.Time } // Delay packets in a bounded pipeline so latency does not serialize every write. func startRecordLinkProxy(t *testing.T, upstream string, mbps int, rtt time.Duration) string { t.Helper() listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatal(err) } ctx, cancel := context.WithCancel(context.Background()) var workers sync.WaitGroup workers.Add(1) go func() { defer workers.Done() left, err := listener.Accept() if err != nil { return } defer left.Close() right, err := (&net.Dialer{}).DialContext(ctx, "tcp", upstream) if err != nil { return } defer right.Close() stopConn := context.AfterFunc(ctx, func() { _ = left.Close(); _ = right.Close() }) defer stopConn() var relays sync.WaitGroup relay := func(dst, src net.Conn) { defer relays.Done() packets := make(chan delayedRecordPacket, 128) var reader sync.WaitGroup reader.Add(1) go func() { defer reader.Done() defer close(packets) buf := make([]byte, 16*1024) for { n, err := src.Read(buf) if n > 0 { packet := delayedRecordPacket{append([]byte(nil), buf[:n]...), time.Now().Add(rtt / 2)} select { case packets <- packet: case <-ctx.Done(): return } } if err != nil { return } } }() defer reader.Wait() defer cancel() var next time.Time for packet := range packets { if next.Before(time.Now()) { next = time.Now() } next = next.Add(time.Duration(len(packet.data)) * time.Second / time.Duration(mbps*1000*1000/8)) ready := packet.ready if next.After(ready) { ready = next } timer := time.NewTimer(time.Until(ready)) select { case <-ctx.Done(): timer.Stop() return case <-timer.C: } if err := writeFullToConn(dst, packet.data); err != nil { return } } } relays.Add(2) go relay(right, left) go relay(left, right) relays.Wait() }() t.Cleanup(func() { cancel(); _ = listener.Close(); workers.Wait() }) return listener.Addr().String() } func TestRecordTCPDelayedBandwidth(t *testing.T) { for _, tc := range []struct { mbps int rtt time.Duration }{{10, 80 * time.Millisecond}, {50, 160 * time.Millisecond}, {100, 80 * time.Millisecond}} { t.Run(fmt.Sprintf("%dMbps-%s", tc.mbps, tc.rtt), func(t *testing.T) { server := NewServer().(*ServerCommon) if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { t.Fatal(err) } const count, size = 512, 2048 handlerDone := make(chan error, 1) server.SetRecordStreamHandler(func(info RecordAcceptInfo) error { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() for i := 0; i < count; i++ { msg, err := info.RecordStream.ReadRecord(ctx) if err != nil { handlerDone <- err return err } want := bytes.Repeat([]byte{byte(i)}, size) binary.BigEndian.PutUint64(want[:8], uint64(i)) if msg.Seq != uint64(i+1) || !bytes.Equal(msg.Payload, want) { err = fmt.Errorf("record %d corrupted", i) handlerDone <- err return err } if err := info.RecordStream.AckRecord(msg.Seq); err != nil { handlerDone <- err return err } } handlerDone <- nil return nil }) if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { t.Fatal(err) } t.Cleanup(func() { _ = server.Stop() }) proxy := startRecordLinkProxy(t, server.listener.Addr().String(), tc.mbps, tc.rtt) client := NewClient().(*ClientCommon) if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { t.Fatal(err) } if err := client.Connect("tcp", proxy); err != nil { t.Fatal(err) } t.Cleanup(func() { _ = client.Stop() }) ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() r, err := client.OpenRecordStream(ctx, RecordOpenOptions{}) if err != nil { t.Fatal(err) } started := time.Now() for i := 0; i < count; i++ { payload := bytes.Repeat([]byte{byte(i)}, size) binary.BigEndian.PutUint64(payload[:8], uint64(i)) if _, err := r.WriteRecord(ctx, payload); err != nil { t.Fatal(err) } } if acked, err := r.Barrier(ctx); err != nil || acked != count { t.Fatalf("barrier ack=%d err=%v", acked, err) } if err := <-handlerDone; err != nil { t.Fatal(err) } if err := r.Close(); err != nil && err != io.EOF { t.Fatal(err) } t.Logf("verified %d records, %d bytes in %s", count, count*size, time.Since(started)) }) } }