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 }