fix(notify): 修复传输生命周期竞态,完善背压与协议边界
- 完善 stream/bulk DataID 分配、预留和双向命名空间,修复并发打开及 dedicated/shared 回退时的 ID 冲突 - 将收发、回复、恢复任务和 sidecar 绑定原始会话与物理连接,防止重连后的旧消息误操作新连接 - 加强 close/reset 身份校验及实例移除检查,修复 dedicated attach 失败、通道引用和资源回收竞态 - 收紧批量发送器停止准入,确保在途入队完成后统一清理请求、缓冲区和等待者 - 修复 record 满队列死锁、取消时序号消耗及关闭竞态,确保关闭有界并返回真实错误 - 增加协商式 record 逻辑半关闭,保留反向 ACK;通过 reset 传递 RecordFailure,避免背压掩盖原始失败原因 - 补齐帧长度、批次数量、序号溢出和未确认窗口校验,提前拒绝超限数据并按字节预算拆批 - 为入站分发增加全局及单连接的条数、字节预算和阻塞背压,关闭时唤醒等待者,消除正常断连日志噪音 - 完善 bulk 窗口释放失败处理与传输诊断,补充并发、重连、背压、协议边界及真实 TCP 回归覆盖
This commit is contained in:
@@ -0,0 +1,317 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user