Files

178 lines
4.7 KiB
Go
Raw Permalink Normal View History

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