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