374 lines
11 KiB
Go
374 lines
11 KiB
Go
|
|
package notify
|
||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"encoding/binary"
|
||
|
|
"errors"
|
||
|
|
"fmt"
|
||
|
|
"io"
|
||
|
|
"sync"
|
||
|
|
"testing"
|
||
|
|
"time"
|
||
|
|
)
|
||
|
|
|
||
|
|
type gatedRecordStream struct {
|
||
|
|
*recordWriteCaptureStream
|
||
|
|
entered chan struct{}
|
||
|
|
proceed chan struct{}
|
||
|
|
first sync.Once
|
||
|
|
unblock sync.Once
|
||
|
|
}
|
||
|
|
|
||
|
|
func newGatedRecordStream() *gatedRecordStream {
|
||
|
|
return &gatedRecordStream{recordWriteCaptureStream: newRecordWriteCaptureStream(), entered: make(chan struct{}), proceed: make(chan struct{})}
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *gatedRecordStream) Write(p []byte) (int, error) {
|
||
|
|
s.first.Do(func() { close(s.entered); <-s.proceed })
|
||
|
|
return s.recordWriteCaptureStream.Write(p)
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *gatedRecordStream) release() { s.unblock.Do(func() { close(s.proceed) }) }
|
||
|
|
func (s *gatedRecordStream) Close() error { s.release(); return s.recordWriteCaptureStream.Close() }
|
||
|
|
func (s *gatedRecordStream) Reset(error) error { return s.Close() }
|
||
|
|
|
||
|
|
type recordWaitContext struct {
|
||
|
|
context.Context
|
||
|
|
waiting chan struct{}
|
||
|
|
once sync.Once
|
||
|
|
}
|
||
|
|
|
||
|
|
func (c *recordWaitContext) Done() <-chan struct{} {
|
||
|
|
c.once.Do(func() { close(c.waiting) })
|
||
|
|
return c.Context.Done()
|
||
|
|
}
|
||
|
|
|
||
|
|
func waitRecordTestSignal(t *testing.T, ch <-chan struct{}) {
|
||
|
|
t.Helper()
|
||
|
|
select {
|
||
|
|
case <-ch:
|
||
|
|
case <-time.After(2 * time.Second):
|
||
|
|
t.Fatal("timed out waiting for record test synchronization")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func blockedRecordForTest(t *testing.T, timeout time.Duration) (*recordStream, *gatedRecordStream) {
|
||
|
|
t.Helper()
|
||
|
|
s := newGatedRecordStream()
|
||
|
|
rs, err := WrapStreamAsRecord(s, RecordOpenOptions{Stream: StreamOpenOptions{WriteTimeout: timeout}})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
r := rs.(*recordStream)
|
||
|
|
t.Cleanup(func() { r.cancel(); _ = s.Close() })
|
||
|
|
if _, err := r.WriteRecord(context.Background(), make([]byte, defaultRecordMaxBatchBytes)); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
waitRecordTestSignal(t, s.entered)
|
||
|
|
return r, s
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordFullQueueResumesAfterSlowWrite(t *testing.T) {
|
||
|
|
r, s := blockedRecordForTest(t, 0)
|
||
|
|
for i := 0; i < cap(r.sendCh); i++ {
|
||
|
|
if _, err := r.WriteRecord(context.Background(), []byte("x")); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
ctx, cancel := context.WithCancel(context.Background())
|
||
|
|
defer cancel()
|
||
|
|
waitCtx := &recordWaitContext{Context: ctx, waiting: make(chan struct{})}
|
||
|
|
done := make(chan error, 1)
|
||
|
|
go func() { _, err := r.WriteRecord(waitCtx, []byte("last")); done <- err }()
|
||
|
|
waitRecordTestSignal(t, waitCtx.waiting)
|
||
|
|
s.release()
|
||
|
|
select {
|
||
|
|
case err := <-done:
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
case <-time.After(time.Second):
|
||
|
|
cancel()
|
||
|
|
<-done
|
||
|
|
t.Fatal("full record queue did not resume after transport recovered")
|
||
|
|
}
|
||
|
|
flushCtx, cancelFlush := context.WithTimeout(context.Background(), time.Second)
|
||
|
|
defer cancelFlush()
|
||
|
|
if err := r.Flush(flushCtx); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordFullQueueCancellationDoesNotConsumeSequence(t *testing.T) {
|
||
|
|
r, s := blockedRecordForTest(t, 0)
|
||
|
|
var last uint64
|
||
|
|
for i := 0; i < cap(r.sendCh); i++ {
|
||
|
|
var err error
|
||
|
|
last, err = r.WriteRecord(context.Background(), []byte("x"))
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
ctx, cancel := context.WithCancel(context.Background())
|
||
|
|
defer cancel()
|
||
|
|
waitCtx := &recordWaitContext{Context: ctx, waiting: make(chan struct{})}
|
||
|
|
done := make(chan error, 1)
|
||
|
|
go func() { _, err := r.WriteRecord(waitCtx, []byte("canceled")); done <- err }()
|
||
|
|
waitRecordTestSignal(t, waitCtx.waiting)
|
||
|
|
cancel()
|
||
|
|
if err := <-done; !errors.Is(err, context.Canceled) {
|
||
|
|
t.Fatalf("canceled write: %v", err)
|
||
|
|
}
|
||
|
|
s.release()
|
||
|
|
retryCtx, cancelRetry := context.WithTimeout(context.Background(), time.Second)
|
||
|
|
defer cancelRetry()
|
||
|
|
seq, err := r.WriteRecord(retryCtx, []byte("retry"))
|
||
|
|
if err != nil || seq != last+1 {
|
||
|
|
t.Fatalf("retry seq=%d err=%v; want %d", seq, err, last+1)
|
||
|
|
}
|
||
|
|
if err := r.Flush(retryCtx); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordCloseBoundsBlockedWrite(t *testing.T) {
|
||
|
|
for _, full := range []bool{false, true} {
|
||
|
|
t.Run(fmt.Sprint(full), func(t *testing.T) {
|
||
|
|
r, _ := blockedRecordForTest(t, 30*time.Millisecond)
|
||
|
|
done := make(chan error, 1)
|
||
|
|
go func() {
|
||
|
|
if full {
|
||
|
|
done <- r.Close()
|
||
|
|
} else {
|
||
|
|
done <- r.CloseWrite()
|
||
|
|
}
|
||
|
|
}()
|
||
|
|
select {
|
||
|
|
case err := <-done:
|
||
|
|
if !errors.Is(err, context.DeadlineExceeded) {
|
||
|
|
t.Fatalf("close error=%v", err)
|
||
|
|
}
|
||
|
|
case <-time.After(time.Second):
|
||
|
|
t.Fatal("close did not interrupt blocked record writer")
|
||
|
|
}
|
||
|
|
if _, err := r.WriteRecord(context.Background(), []byte("late")); err == nil {
|
||
|
|
t.Fatal("write succeeded after close")
|
||
|
|
}
|
||
|
|
waitRecordTestSignal(t, r.Context().Done())
|
||
|
|
})
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordCloseRejectsLateWritesAndIsIdempotent(t *testing.T) {
|
||
|
|
for _, full := range []bool{false, true} {
|
||
|
|
t.Run(fmt.Sprint(full), func(t *testing.T) {
|
||
|
|
s := newRecordWriteCaptureStream()
|
||
|
|
r, err := WrapStreamAsRecord(s, RecordOpenOptions{})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
t.Cleanup(func() { _ = r.Close() })
|
||
|
|
closeStream := r.CloseWrite
|
||
|
|
if full {
|
||
|
|
closeStream = r.Close
|
||
|
|
}
|
||
|
|
if err := closeStream(); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if err := closeStream(); err != nil {
|
||
|
|
t.Fatalf("repeated close: %v", err)
|
||
|
|
}
|
||
|
|
for i := 0; i < 32; i++ {
|
||
|
|
seq, err := r.WriteRecord(context.Background(), []byte("late"))
|
||
|
|
if err == nil || seq != 0 {
|
||
|
|
t.Fatalf("closed write seq=%d err=%v", seq, err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordCanceledContextNeverAdmitsWrite(t *testing.T) {
|
||
|
|
s := newRecordWriteCaptureStream()
|
||
|
|
r, err := WrapStreamAsRecord(s, RecordOpenOptions{})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer r.Close()
|
||
|
|
ctx, cancel := context.WithCancel(context.Background())
|
||
|
|
cancel()
|
||
|
|
if seq, err := r.WriteRecord(ctx, []byte("canceled")); seq != 0 || !errors.Is(err, context.Canceled) {
|
||
|
|
t.Fatalf("seq=%d err=%v", seq, err)
|
||
|
|
}
|
||
|
|
if seq, err := r.WriteRecord(context.Background(), []byte("first")); seq != 1 || err != nil {
|
||
|
|
t.Fatalf("seq=%d err=%v", seq, err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordLegacyLargeBatchesAndOptions(t *testing.T) {
|
||
|
|
const maxWireRecordCount = 1<<16 - 1
|
||
|
|
for _, count := range []int{65, 512, 2048, maxWireRecordCount} {
|
||
|
|
for _, v := range []byte{recordFrameVersionV1, recordFrameVersionV2} {
|
||
|
|
t.Run(fmt.Sprintf("v%d/%d", v, count), func(t *testing.T) {
|
||
|
|
frame := makeRecordBatchFrameHeaderForBoundsTest(v, count)
|
||
|
|
for i := 0; i < count; i++ {
|
||
|
|
frame = append(frame, 0, 0, 0, 1, 'x')
|
||
|
|
}
|
||
|
|
decoded, err := decodeRecordFrame(frame)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if len(decoded.Batch) != count || decoded.Batch[count-1].Seq != uint64(count) {
|
||
|
|
t.Fatal("batch truncated")
|
||
|
|
}
|
||
|
|
reencoded, err := encodeRecordBatchFrame(decoded.Batch, 0, v == recordFrameVersionV2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if len(reencoded) != len(frame) {
|
||
|
|
t.Fatal("wire size changed")
|
||
|
|
}
|
||
|
|
opt := normalizeRecordOpenOptions(RecordOpenOptions{MaxBatchRecords: count})
|
||
|
|
if opt.MaxBatchRecords != count {
|
||
|
|
t.Fatalf("batch count reduced to %d", opt.MaxBatchRecords)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
}
|
||
|
|
tooMany := make([]recordOutboundMessage, maxWireRecordCount+1)
|
||
|
|
if _, err := encodeRecordBatchFrame(tooMany, 0, false); !errors.Is(err, errRecordFrameInvalid) {
|
||
|
|
t.Fatalf("oversized count: %v", err)
|
||
|
|
}
|
||
|
|
for _, v := range []byte{recordFrameVersionV1, recordFrameVersionV2} {
|
||
|
|
frame := makeRecordBatchFrameHeaderForBoundsTest(v, 2)
|
||
|
|
binary.BigEndian.PutUint64(frame[10:18], ^uint64(0))
|
||
|
|
frame = append(frame, make([]byte, 8)...)
|
||
|
|
if _, err := decodeRecordFrame(frame); !errors.Is(err, errRecordFrameInvalid) {
|
||
|
|
t.Fatalf("wrapping sequence: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
type failingRecordWriteStream struct{ *recordWriteCaptureStream }
|
||
|
|
|
||
|
|
func (s *failingRecordWriteStream) Write([]byte) (int, error) { return 0, io.ErrUnexpectedEOF }
|
||
|
|
|
||
|
|
func TestRecordFlushWriteFailureIsTerminal(t *testing.T) {
|
||
|
|
s := &failingRecordWriteStream{newRecordWriteCaptureStream()}
|
||
|
|
rs, err := WrapStreamAsRecord(s, RecordOpenOptions{MaxBatchDelay: time.Hour})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
r := rs.(*recordStream)
|
||
|
|
t.Cleanup(func() { r.cancel(); _ = s.Close() })
|
||
|
|
if _, err := r.WriteRecord(context.Background(), []byte("fail")); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||
|
|
defer cancel()
|
||
|
|
if err := r.Flush(ctx); !errors.Is(err, io.ErrUnexpectedEOF) {
|
||
|
|
t.Fatalf("flush: %v", err)
|
||
|
|
}
|
||
|
|
waitRecordTestSignal(t, r.Context().Done())
|
||
|
|
if _, err := r.WriteRecord(ctx, []byte("late")); !errors.Is(err, io.ErrUnexpectedEOF) {
|
||
|
|
t.Fatalf("late write: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordCloseAbortsNativeStreamAndReleasesRuntime(t *testing.T) {
|
||
|
|
runtime := newStreamRuntime("record-close")
|
||
|
|
entered, finished, resetSent := make(chan struct{}), make(chan struct{}), make(chan struct{})
|
||
|
|
s := newStreamHandle(context.Background(), runtime, clientFileScope(), StreamOpenRequest{StreamID: "record-close"}, 0, nil, nil, 0, nil,
|
||
|
|
func(context.Context, *streamHandle, string) error { close(resetSent); return nil },
|
||
|
|
func(ctx context.Context, _ *streamHandle, _ []byte) error {
|
||
|
|
close(entered)
|
||
|
|
<-ctx.Done()
|
||
|
|
close(finished)
|
||
|
|
return ctx.Err()
|
||
|
|
}, defaultStreamConfig())
|
||
|
|
if err := runtime.register(clientFileScope(), s); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
t.Cleanup(func() { s.markReset(io.ErrClosedPipe) })
|
||
|
|
rs, err := WrapStreamAsRecord(s, RecordOpenOptions{Stream: StreamOpenOptions{WriteTimeout: 30 * time.Millisecond}})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
r := rs.(*recordStream)
|
||
|
|
if _, err := r.WriteRecord(context.Background(), make([]byte, defaultRecordMaxBatchBytes)); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
waitRecordTestSignal(t, entered)
|
||
|
|
if err := r.Close(); !errors.Is(err, context.DeadlineExceeded) {
|
||
|
|
t.Fatalf("close: %v", err)
|
||
|
|
}
|
||
|
|
waitRecordTestSignal(t, finished)
|
||
|
|
waitRecordTestSignal(t, resetSent)
|
||
|
|
waitRecordTestSignal(t, r.readerCh)
|
||
|
|
waitRecordTestSignal(t, r.writerCh)
|
||
|
|
if _, ok := runtime.lookup(clientFileScope(), s.ID()); ok {
|
||
|
|
t.Fatal("closed record retained underlying stream")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordCloseSealsAdmissionBeforeDrain(t *testing.T) {
|
||
|
|
r, s := blockedRecordForTest(t, time.Second)
|
||
|
|
done := make(chan error, 1)
|
||
|
|
go func() { done <- r.CloseWrite() }()
|
||
|
|
deadline := time.Now().Add(time.Second)
|
||
|
|
for {
|
||
|
|
r.mu.Lock()
|
||
|
|
sealed := r.outboundClosed
|
||
|
|
r.mu.Unlock()
|
||
|
|
if sealed {
|
||
|
|
break
|
||
|
|
}
|
||
|
|
if time.Now().After(deadline) {
|
||
|
|
t.Fatal("close did not seal admission")
|
||
|
|
}
|
||
|
|
time.Sleep(time.Millisecond)
|
||
|
|
}
|
||
|
|
if _, err := r.WriteRecord(context.Background(), []byte("late")); !errors.Is(err, errRecordWriteClosed) {
|
||
|
|
t.Fatalf("late write: %v", err)
|
||
|
|
}
|
||
|
|
s.release()
|
||
|
|
if err := <-done; err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if len(s.Bytes()) == 0 {
|
||
|
|
t.Fatal("close lost accepted data")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
type cancelOnCloseRecordStream struct {
|
||
|
|
*recordWriteCaptureStream
|
||
|
|
ctx context.Context
|
||
|
|
cancel context.CancelFunc
|
||
|
|
closing, finish chan struct{}
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *cancelOnCloseRecordStream) Context() context.Context { return s.ctx }
|
||
|
|
func (s *cancelOnCloseRecordStream) Close() error {
|
||
|
|
s.cancel()
|
||
|
|
close(s.closing)
|
||
|
|
<-s.finish
|
||
|
|
return s.recordWriteCaptureStream.Close()
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRecordCloseWaitsForUnderlyingCloseResult(t *testing.T) {
|
||
|
|
ctx, cancel := context.WithCancel(context.Background())
|
||
|
|
defer cancel()
|
||
|
|
s := &cancelOnCloseRecordStream{recordWriteCaptureStream: newRecordWriteCaptureStream(), ctx: ctx, cancel: cancel, closing: make(chan struct{}), finish: make(chan struct{})}
|
||
|
|
r, err := WrapStreamAsRecord(s, RecordOpenOptions{})
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
done := make(chan error, 1)
|
||
|
|
go func() { done <- r.Close() }()
|
||
|
|
waitRecordTestSignal(t, s.closing)
|
||
|
|
close(s.finish)
|
||
|
|
if err := <-done; err != nil {
|
||
|
|
t.Fatalf("successful close reported cancellation: %v", err)
|
||
|
|
}
|
||
|
|
}
|