Files

374 lines
11 KiB
Go
Raw Permalink Normal View History

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