318 lines
13 KiB
Go
318 lines
13 KiB
Go
|
|
package notify
|
||
|
|
|
||
|
|
import (
|
||
|
|
"bytes"
|
||
|
|
"context"
|
||
|
|
"encoding/binary"
|
||
|
|
"errors"
|
||
|
|
"io"
|
||
|
|
"net"
|
||
|
|
"sync"
|
||
|
|
"testing"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"b612.me/stario"
|
||
|
|
)
|
||
|
|
|
||
|
|
type recordWriteCaptureStream struct {
|
||
|
|
mu sync.Mutex
|
||
|
|
buf bytes.Buffer
|
||
|
|
readDone chan struct{}
|
||
|
|
close sync.Once
|
||
|
|
}
|
||
|
|
|
||
|
|
func newRecordWriteCaptureStream() *recordWriteCaptureStream {
|
||
|
|
return &recordWriteCaptureStream{
|
||
|
|
readDone: make(chan struct{}),
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *recordWriteCaptureStream) Read([]byte) (int, error) {
|
||
|
|
<-s.readDone
|
||
|
|
return 0, io.EOF
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *recordWriteCaptureStream) Write(p []byte) (int, error) {
|
||
|
|
s.mu.Lock()
|
||
|
|
defer s.mu.Unlock()
|
||
|
|
return s.buf.Write(p)
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *recordWriteCaptureStream) Close() error {
|
||
|
|
s.close.Do(func() {
|
||
|
|
close(s.readDone)
|
||
|
|
})
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *recordWriteCaptureStream) ID() string { return "record-capture" }
|
||
|
|
func (s *recordWriteCaptureStream) Channel() StreamChannel { return StreamRecordChannel }
|
||
|
|
func (s *recordWriteCaptureStream) Metadata() StreamMetadata { return nil }
|
||
|
|
func (s *recordWriteCaptureStream) Context() context.Context { return context.Background() }
|
||
|
|
func (s *recordWriteCaptureStream) LogicalConn() *LogicalConn { return nil }
|
||
|
|
func (s *recordWriteCaptureStream) TransportConn() *TransportConn { return nil }
|
||
|
|
func (s *recordWriteCaptureStream) TransportGeneration() uint64 { return 0 }
|
||
|
|
func (s *recordWriteCaptureStream) LocalAddr() net.Addr { return nil }
|
||
|
|
func (s *recordWriteCaptureStream) RemoteAddr() net.Addr { return nil }
|
||
|
|
func (s *recordWriteCaptureStream) CloseWrite() error { return nil }
|
||
|
|
func (s *recordWriteCaptureStream) Reset(error) error { return s.Close() }
|
||
|
|
func (s *recordWriteCaptureStream) SetDeadline(time.Time) error { return nil }
|
||
|
|
func (s *recordWriteCaptureStream) SetReadDeadline(time.Time) error {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
func (s *recordWriteCaptureStream) SetWriteDeadline(time.Time) error {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *recordWriteCaptureStream) Bytes() []byte {
|
||
|
|
s.mu.Lock()
|
||
|
|
defer s.mu.Unlock()
|
||
|
|
return append([]byte(nil), s.buf.Bytes()...)
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestDedicatedRecordRejectsOversizedPayloadLength(t *testing.T) {
|
||
|
|
conn := &shortWriteBulkRecordConn{maxPerWrite: bulkDedicatedRecordMaxBytes + bulkDedicatedRecordHeaderLen + 1}
|
||
|
|
err := writeBulkDedicatedRecordWithDeadline(conn, make([]byte, bulkDedicatedRecordMaxBytes+1), time.Time{})
|
||
|
|
if !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("writeBulkDedicatedRecordWithDeadline error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
if got := conn.buf.Len(); got != 0 {
|
||
|
|
t.Fatalf("oversized dedicated record wrote %d bytes, want 0", got)
|
||
|
|
}
|
||
|
|
|
||
|
|
header := make([]byte, bulkDedicatedRecordHeaderLen)
|
||
|
|
copy(header[:4], bulkDedicatedRecordMagic)
|
||
|
|
binary.BigEndian.PutUint32(header[4:8], uint32(bulkDedicatedRecordMaxBytes+1))
|
||
|
|
_, release, err := readBulkDedicatedRecordPooled(newBulkAttachScriptConn(header))
|
||
|
|
if release != nil {
|
||
|
|
release()
|
||
|
|
t.Fatal("oversized dedicated record returned release callback")
|
||
|
|
}
|
||
|
|
if !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("readBulkDedicatedRecordPooled error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestDirectSignalFrameRejectsOversizedPayloadLength(t *testing.T) {
|
||
|
|
header := stario.NewQueue().BuildHeader(uint32(transportFrameMaxPayloadBytes + 1))
|
||
|
|
_, err := readDirectSignalFramePayload(newBulkAttachScriptConn(header))
|
||
|
|
if !errors.Is(err, stario.ErrQueueMessageTooLarge) {
|
||
|
|
t.Fatalf("readDirectSignalFramePayload error = %v, want %v", err, stario.ErrQueueMessageTooLarge)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestTransferFrameRejectsOversizedPayloadLength(t *testing.T) {
|
||
|
|
stream := &transferWriteCountStream{}
|
||
|
|
var header [transferFrameHeaderSize]byte
|
||
|
|
binary.BigEndian.PutUint32(header[:], uint32(transferFrameMaxPayloadBytes+1))
|
||
|
|
if _, err := stream.buf.Write(header[:]); err != nil {
|
||
|
|
t.Fatalf("seed transfer frame header failed: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
_, err := readTransferFrame(stream)
|
||
|
|
if !errors.Is(err, errTransferFrameTooLarge) {
|
||
|
|
t.Fatalf("readTransferFrame error = %v, want %v", err, errTransferFrameTooLarge)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestDedicatedBatchDecodersRejectOversizedWireCounts(t *testing.T) {
|
||
|
|
tooManyItems := make([]bulkDedicatedSendRequest, bulkDedicatedBatchMaxItems+1)
|
||
|
|
for i := range tooManyItems {
|
||
|
|
tooManyItems[i] = bulkDedicatedSendRequest{Type: bulkFastPayloadTypeData, Seq: uint64(i + 1)}
|
||
|
|
}
|
||
|
|
if _, err := encodeBulkDedicatedBatchPlain(1, tooManyItems); !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("encodeBulkDedicatedBatchPlain oversized item count error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
tooManyGroups := make([]bulkDedicatedOutboundBatch, bulkDedicatedBatchMaxItems+1)
|
||
|
|
for i := range tooManyGroups {
|
||
|
|
tooManyGroups[i] = bulkDedicatedOutboundBatch{
|
||
|
|
DataID: uint64(i + 1),
|
||
|
|
Items: []bulkDedicatedSendRequest{{
|
||
|
|
Type: bulkFastPayloadTypeData,
|
||
|
|
Seq: 1,
|
||
|
|
}},
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if _, err := encodeBulkDedicatedBatchesPlain(tooManyGroups); !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("encodeBulkDedicatedBatchesPlain oversized group count error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
batch := make([]byte, bulkDedicatedBatchHeaderLen)
|
||
|
|
copy(batch[:4], bulkDedicatedBatchMagic)
|
||
|
|
batch[4] = bulkDedicatedBatchVersion
|
||
|
|
binary.BigEndian.PutUint64(batch[8:16], 1)
|
||
|
|
binary.BigEndian.PutUint32(batch[16:20], uint32(bulkDedicatedBatchMaxItems+1))
|
||
|
|
if _, _, matched, err := decodeBulkDedicatedBatchPlain(batch); !matched || !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("decodeBulkDedicatedBatchPlain matched=%v error=%v, want matched invalid", matched, err)
|
||
|
|
}
|
||
|
|
if err := walkDedicatedBulkInboundBatchPlain(batch, func(uint64, bulkDedicatedBatchItem) error {
|
||
|
|
t.Fatal("visit should not be called for oversized batch count")
|
||
|
|
return nil
|
||
|
|
}); !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("walkDedicatedBulkInboundBatchPlain error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
superGroups := make([]byte, bulkDedicatedSuperBatchHeaderLen)
|
||
|
|
copy(superGroups[:4], bulkDedicatedSuperBatchMagic)
|
||
|
|
superGroups[4] = bulkDedicatedSuperBatchVersion
|
||
|
|
binary.BigEndian.PutUint32(superGroups[8:12], uint32(bulkDedicatedBatchMaxItems+1))
|
||
|
|
if _, matched, err := decodeBulkDedicatedSuperBatchPlain(superGroups); !matched || !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("decodeBulkDedicatedSuperBatchPlain groups matched=%v error=%v, want matched invalid", matched, err)
|
||
|
|
}
|
||
|
|
if err := walkDedicatedBulkInboundSuperBatchPlain(superGroups, func(uint64, bulkDedicatedBatchItem) error {
|
||
|
|
t.Fatal("visit should not be called for oversized super-batch group count")
|
||
|
|
return nil
|
||
|
|
}); !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("walkDedicatedBulkInboundSuperBatchPlain groups error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
superItems := make([]byte, bulkDedicatedSuperBatchHeaderLen+bulkDedicatedSuperBatchGroupHeaderLen)
|
||
|
|
copy(superItems[:4], bulkDedicatedSuperBatchMagic)
|
||
|
|
superItems[4] = bulkDedicatedSuperBatchVersion
|
||
|
|
binary.BigEndian.PutUint32(superItems[8:12], 1)
|
||
|
|
binary.BigEndian.PutUint64(superItems[12:20], 1)
|
||
|
|
binary.BigEndian.PutUint32(superItems[20:24], uint32(bulkDedicatedBatchMaxItems+1))
|
||
|
|
if _, matched, err := decodeBulkDedicatedSuperBatchPlain(superItems); !matched || !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("decodeBulkDedicatedSuperBatchPlain items matched=%v error=%v, want matched invalid", matched, err)
|
||
|
|
}
|
||
|
|
if err := walkDedicatedBulkInboundSuperBatchPlain(superItems, func(uint64, bulkDedicatedBatchItem) error {
|
||
|
|
t.Fatal("visit should not be called for oversized super-batch item count")
|
||
|
|
return nil
|
||
|
|
}); !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("walkDedicatedBulkInboundSuperBatchPlain items error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestSharedFastBatchDecodersRejectOversizedWireCounts(t *testing.T) {
|
||
|
|
tooManyBulkFrames := make([]bulkFastFrame, bulkFastBatchMaxItems+1)
|
||
|
|
for i := range tooManyBulkFrames {
|
||
|
|
tooManyBulkFrames[i] = bulkFastFrame{Type: bulkFastPayloadTypeData, DataID: 1, Seq: uint64(i + 1)}
|
||
|
|
}
|
||
|
|
if _, err := encodeBulkFastBatchPlain(tooManyBulkFrames); !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("encodeBulkFastBatchPlain oversized count error = %v, want %v", err, errBulkFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
bulkBatch := make([]byte, bulkFastBatchHeaderLen)
|
||
|
|
copy(bulkBatch[:4], bulkFastBatchMagic)
|
||
|
|
bulkBatch[4] = bulkFastBatchVersion
|
||
|
|
binary.BigEndian.PutUint32(bulkBatch[8:12], uint32(bulkFastBatchMaxItems+1))
|
||
|
|
if matched, err := walkBulkFastBatchPlain(bulkBatch, func(bulkFastFrame) error {
|
||
|
|
t.Fatal("bulk batch visitor should not be called for oversized count")
|
||
|
|
return nil
|
||
|
|
}); !matched || !errors.Is(err, errBulkFastPayloadInvalid) {
|
||
|
|
t.Fatalf("walkBulkFastBatchPlain matched=%v error=%v, want matched invalid", matched, err)
|
||
|
|
}
|
||
|
|
|
||
|
|
tooManyStreamFrames := make([]streamFastDataFrame, streamFastBatchMaxItems+1)
|
||
|
|
for i := range tooManyStreamFrames {
|
||
|
|
tooManyStreamFrames[i] = streamFastDataFrame{DataID: 1, Seq: uint64(i + 1)}
|
||
|
|
}
|
||
|
|
if _, err := encodeStreamFastBatchPlain(tooManyStreamFrames); !errors.Is(err, errStreamFastPayloadInvalid) {
|
||
|
|
t.Fatalf("encodeStreamFastBatchPlain oversized count error = %v, want %v", err, errStreamFastPayloadInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
streamBatch := make([]byte, streamFastBatchHeaderLen)
|
||
|
|
copy(streamBatch[:4], streamFastBatchMagic)
|
||
|
|
streamBatch[4] = streamFastBatchVersion
|
||
|
|
binary.BigEndian.PutUint32(streamBatch[8:12], uint32(streamFastBatchMaxItems+1))
|
||
|
|
if matched, err := walkStreamFastBatchPlain(streamBatch, func(streamFastDataFrame) error {
|
||
|
|
t.Fatal("stream batch visitor should not be called for oversized count")
|
||
|
|
return nil
|
||
|
|
}); !matched || !errors.Is(err, errStreamFastPayloadInvalid) {
|
||
|
|
t.Fatalf("walkStreamFastBatchPlain matched=%v error=%v, want matched invalid", matched, err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordStreamRejectsPayloadLargerThanUnackedWindow(t *testing.T) {
|
||
|
|
record := &recordStream{
|
||
|
|
cfg: recordConfig{
|
||
|
|
MaxUnackedBytes: 4,
|
||
|
|
},
|
||
|
|
}
|
||
|
|
_, err := record.WriteRecord(context.Background(), []byte("12345"))
|
||
|
|
if !errors.Is(err, errRecordPayloadTooLarge) {
|
||
|
|
t.Fatalf("WriteRecord error = %v, want %v", err, errRecordPayloadTooLarge)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordOptionsCapBatchCountsToWireLimit(t *testing.T) {
|
||
|
|
opt := normalizeRecordOpenOptions(RecordOpenOptions{
|
||
|
|
MaxBatchRecords: recordMaxBatchRecords + 100,
|
||
|
|
MaxBatchBytes: transferFrameMaxPayloadBytes * 2,
|
||
|
|
MaxUnackedBytes: transferFrameMaxPayloadBytes * 2,
|
||
|
|
})
|
||
|
|
if got, want := opt.MaxBatchRecords, recordMaxBatchRecords; got != want {
|
||
|
|
t.Fatalf("MaxBatchRecords = %d, want %d", got, want)
|
||
|
|
}
|
||
|
|
if got, want := opt.MaxBatchBytes, recordMaxBatchPayloadBytes; got != want {
|
||
|
|
t.Fatalf("MaxBatchBytes = %d, want %d", got, want)
|
||
|
|
}
|
||
|
|
if got, want := opt.MaxUnackedBytes, transferFrameMaxPayloadBytes*2; got != want {
|
||
|
|
t.Fatalf("MaxUnackedBytes = %d, want %d", got, want)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordWriterFlushesBeforeAppendingPastBatchByteLimit(t *testing.T) {
|
||
|
|
stream := newRecordWriteCaptureStream()
|
||
|
|
record, err := WrapStreamAsRecord(stream, RecordOpenOptions{
|
||
|
|
MaxBatchRecords: defaultRecordMaxBatchRecords,
|
||
|
|
MaxBatchBytes: 10,
|
||
|
|
MaxBatchDelay: time.Hour,
|
||
|
|
MaxUnackedRecords: 16,
|
||
|
|
MaxUnackedBytes: 1024,
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("WrapStreamAsRecord failed: %v", err)
|
||
|
|
}
|
||
|
|
defer record.Close()
|
||
|
|
|
||
|
|
if _, err := record.WriteRecord(context.Background(), []byte("123456")); err != nil {
|
||
|
|
t.Fatalf("first WriteRecord failed: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := record.WriteRecord(context.Background(), []byte("abcdef")); err != nil {
|
||
|
|
t.Fatalf("second WriteRecord failed: %v", err)
|
||
|
|
}
|
||
|
|
flushCtx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||
|
|
defer cancel()
|
||
|
|
if err := record.Flush(flushCtx); err != nil {
|
||
|
|
t.Fatalf("Flush failed: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
if got, want := countTransferFrames(stream.Bytes()), 2; got != want {
|
||
|
|
t.Fatalf("transfer frame count = %d, want %d", got, want)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordBatchDecodersRejectOversizedWireCountsAndLengths(t *testing.T) {
|
||
|
|
v1 := makeRecordBatchFrameHeaderForBoundsTest(recordFrameVersionV1, defaultRecordMaxBatchRecords+1)
|
||
|
|
if _, err := decodeRecordFrame(v1); !errors.Is(err, errRecordFrameInvalid) {
|
||
|
|
t.Fatalf("decodeRecordFrame oversized v1 count error = %v, want %v", err, errRecordFrameInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
v2 := makeRecordBatchFrameHeaderForBoundsTest(recordFrameVersionV2, defaultRecordMaxBatchRecords+1)
|
||
|
|
if _, err := decodeRecordFrame(v2); !errors.Is(err, errRecordFrameInvalid) {
|
||
|
|
t.Fatalf("decodeRecordFrame oversized v2 count error = %v, want %v", err, errRecordFrameInvalid)
|
||
|
|
}
|
||
|
|
|
||
|
|
v2LongItem := makeRecordBatchFrameHeaderForBoundsTest(recordFrameVersionV2, 1)
|
||
|
|
v2LongItem = append(v2LongItem, 0xff, 0xff, 0xff, 0xff)
|
||
|
|
if _, err := decodeRecordFrame(v2LongItem); !errors.Is(err, errRecordFrameInvalid) {
|
||
|
|
t.Fatalf("decodeRecordFrame oversized v2 item length error = %v, want %v", err, errRecordFrameInvalid)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func makeRecordBatchFrameHeaderForBoundsTest(version uint8, count int) []byte {
|
||
|
|
headerSize := recordBatchHeaderV1Size
|
||
|
|
if version == recordFrameVersionV2 {
|
||
|
|
headerSize = recordBatchHeaderV2Size
|
||
|
|
}
|
||
|
|
frame := make([]byte, recordFrameHeaderSize+headerSize)
|
||
|
|
copy(frame[:4], recordFrameMagic)
|
||
|
|
frame[4] = version
|
||
|
|
frame[5] = recordFrameTypeBatch
|
||
|
|
binary.BigEndian.PutUint16(frame[8:10], uint16(count))
|
||
|
|
binary.BigEndian.PutUint64(frame[10:18], 1)
|
||
|
|
return frame
|
||
|
|
}
|