Files
notify/record_protocol_lifecycle_test.go
T

319 lines
11 KiB
Go
Raw Normal View History

package notify
import (
"bytes"
"context"
"errors"
"io"
"sync"
"sync/atomic"
"testing"
"time"
)
func TestRegressionRecordResetSendFailureReleasesNativeStream(t *testing.T) {
runtime := newStreamRuntime("review-reset")
sendErr := errors.New("injected reset send failure")
s := newStreamHandle(context.Background(), runtime, clientFileScope(), StreamOpenRequest{StreamID: "review-reset"}, 0, nil, nil, 0, nil,
func(context.Context, *streamHandle, string) error { return sendErr }, nil, defaultStreamConfig())
if err := runtime.register(clientFileScope(), s); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { s.markReset(io.ErrClosedPipe) })
record, err := WrapStreamAsRecord(s, RecordOpenOptions{})
if err != nil {
t.Fatal(err)
}
err = record.Reset(errors.New("application aborted"))
if err != nil && !errors.Is(err, sendErr) {
t.Fatal(err)
}
select {
case <-s.Context().Done():
case <-time.After(100 * time.Millisecond):
_, retained := runtime.lookup(clientFileScope(), s.ID())
t.Fatalf("Reset returned %v but native context is live, runtime retained=%v", err, retained)
}
select {
case <-record.(*recordStream).readerCh:
case <-time.After(time.Second):
t.Fatal("record reader was not released")
}
}
func TestRegressionRecordHalfCloseKeepsResponseAcknowledgements(t *testing.T) {
server := NewServer().(*ServerCommon)
if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil {
t.Fatal(err)
}
accepted := make(chan RecordStream, 1)
server.SetRecordStreamHandler(func(info RecordAcceptInfo) error { accepted <- info.RecordStream; return nil })
if err := server.Listen("tcp", "127.0.0.1:0"); err != nil {
t.Fatal(err)
}
defer server.Stop()
client := NewClient().(*ClientCommon)
if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil {
t.Fatal(err)
}
if err := client.Connect("tcp", server.listener.Addr().String()); err != nil {
t.Fatal(err)
}
defer client.Stop()
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
local, err := client.OpenRecordStream(ctx, RecordOpenOptions{})
if err != nil {
t.Fatal(err)
}
var remote RecordStream
select {
case remote = <-accepted:
case <-ctx.Done():
t.Fatal(ctx.Err())
}
for i := 0; i < 3; i++ {
if _, err := local.WriteRecord(ctx, []byte{byte(i)}); err != nil {
t.Fatal(err)
}
}
if err := local.CloseWrite(); err != nil {
t.Fatal(err)
}
for i := 0; i < 3; i++ {
msg, err := remote.ReadRecord(ctx)
if err != nil || !bytes.Equal(msg.Payload, []byte{byte(i)}) {
t.Fatalf("queued request %d: %v %+v", i, err, msg)
}
if err := remote.AckRecord(msg.Seq); err != nil {
t.Fatal(err)
}
}
if _, err := remote.ReadRecord(ctx); !errors.Is(err, io.EOF) {
t.Fatalf("request half-close EOF: %v", err)
}
if _, err := local.BarrierTo(ctx, 3); err != nil {
t.Fatalf("request ACK after FIN: %v", err)
}
seq, err := remote.WriteRecord(ctx, []byte("response after request EOF"))
if err != nil {
t.Fatal(err)
}
if err := remote.Flush(ctx); err != nil {
t.Fatal(err)
}
msg, err := local.ReadRecord(ctx)
if err != nil {
t.Fatal(err)
}
if err := local.AckRecord(msg.Seq); err != nil {
t.Fatal(err)
}
if _, err := remote.BarrierTo(ctx, seq); err != nil {
t.Fatalf("response was received and applied, but its barrier failed after peer CloseWrite: %v", err)
}
}
func TestRecordHalfCloseNegotiation(t *testing.T) {
request := advertiseRecordStreamOpenMetadata(StreamMetadata{recordStreamMetadataUseHalfCloseKey: "1"})
if recordStreamUseHalfClose(request) {
t.Fatal("request enabled half-close before peer negotiation")
}
for _, supported := range []bool{false, true} {
meta := StreamMetadata{recordStreamMetadataUseHalfCloseKey: "1"}
if supported {
meta[recordStreamMetadataCapHalfCloseKey] = "1"
}
accepted, response := negotiateRecordStreamOpenMetadata(StreamRecordChannel, meta)
if recordStreamUseHalfClose(accepted) != supported || recordStreamUseHalfClose(response) != supported {
t.Fatalf("half-close negotiation with supported=%v: %v %v", supported, accepted, response)
}
}
}
func TestRecordLegacyHalfCloseUsesNativeClose(t *testing.T) {
s := newStreamHandle(context.Background(), newStreamRuntime("legacy-fin"), clientFileScope(), StreamOpenRequest{StreamID: "legacy-fin"}, 0, nil, nil, 0, nil, nil, nil, defaultStreamConfig())
t.Cleanup(func() { s.markReset(io.ErrClosedPipe) })
record, err := WrapStreamAsRecord(s, RecordOpenOptions{})
if err != nil {
t.Fatal(err)
}
if err := record.CloseWrite(); err != nil {
t.Fatal(err)
}
if !s.localClosedSnapshot() {
t.Fatal("unnegotiated peer received logical FIN instead of native close")
}
}
func TestRecordInvalidFINAndPostFINDataAbort(t *testing.T) {
batch, err := encodeRecordBatchFrame([]recordOutboundMessage{{Seq: 1, Payload: []byte("after EOF")}}, 0, false)
if err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
negotiated bool
frames [][]byte
}{
{"unnegotiated", false, [][]byte{encodeRecordFINFrame(0)}},
{"wrong-sequence", true, [][]byte{encodeRecordFINFrame(1)}},
{"duplicate", true, [][]byte{encodeRecordFINFrame(0), encodeRecordFINFrame(0)}},
{"data-after-fin", true, [][]byte{encodeRecordFINFrame(0), batch}},
{"truncated-fin", true, [][]byte{encodeRecordFINFrame(0)[:10]}},
} {
t.Run(tc.name, func(t *testing.T) {
metadata := StreamMetadata{}
if tc.negotiated {
metadata[recordStreamMetadataUseHalfCloseKey] = "1"
}
s := newStreamHandle(context.Background(), newStreamRuntime("fin"), clientFileScope(), StreamOpenRequest{StreamID: "fin", Metadata: metadata}, 0, nil, nil, 0, nil, nil,
func(context.Context, *streamHandle, []byte) error { return nil }, defaultStreamConfig())
t.Cleanup(func() { s.markReset(io.ErrClosedPipe) })
record, err := WrapStreamAsRecord(s, RecordOpenOptions{})
if err != nil {
t.Fatal(err)
}
for _, payload := range tc.frames {
if err := s.pushChunk(buildTransferFrame(payload)); err != nil {
t.Fatal(err)
}
}
select {
case <-record.(*recordStream).readerCh:
case <-time.After(time.Second):
t.Fatal("invalid FIN reader stuck")
}
if s.Context().Err() == nil {
t.Fatal("protocol failure retained native stream")
}
if _, err := record.WriteRecord(context.Background(), []byte("unexpected")); err == nil {
t.Fatal("protocol failure accepted a new write")
}
})
}
}
func TestRecordFailureNotificationIsBounded(t *testing.T) {
s := newStreamHandle(context.Background(), newStreamRuntime("blocked-failure"), clientFileScope(), StreamOpenRequest{StreamID: "blocked-failure"}, 0, nil, nil, 0, nil, nil,
func(_ context.Context, stream *streamHandle, _ []byte) error {
<-stream.Context().Done()
return io.ErrClosedPipe
}, defaultStreamConfig())
t.Cleanup(func() { s.markReset(io.ErrClosedPipe) })
record, err := WrapStreamAsRecord(s, RecordOpenOptions{Stream: StreamOpenOptions{WriteTimeout: 30 * time.Millisecond}})
if err != nil {
t.Fatal(err)
}
started := time.Now()
if err := record.FailRecord(1, RecordFailure{Message: "apply failed"}); !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("bounded failure error: %v", err)
}
if elapsed := time.Since(started); elapsed > time.Second {
t.Fatalf("failure took %v", elapsed)
}
if s.Context().Err() == nil {
t.Fatal("failure retained native stream")
}
select {
case <-record.(*recordStream).readerCh:
case <-time.After(time.Second):
t.Fatal("failure retained reader")
}
}
func TestRecordWriterFailureReleasesNativeStream(t *testing.T) {
errWrite := errors.New("injected write error")
s := newStreamHandle(context.Background(), newStreamRuntime("writer-failure"), clientFileScope(), StreamOpenRequest{StreamID: "writer-failure"}, 0, nil, nil, 0, nil, nil,
func(context.Context, *streamHandle, []byte) error { return errWrite }, defaultStreamConfig())
t.Cleanup(func() { s.markReset(io.ErrClosedPipe) })
record, err := WrapStreamAsRecord(s, RecordOpenOptions{MaxBatchRecords: 1})
if err != nil {
t.Fatal(err)
}
if _, err := record.WriteRecord(context.Background(), []byte("request")); err != nil {
t.Fatal(err)
}
select {
case <-record.(*recordStream).writerCh:
case <-time.After(time.Second):
t.Fatal("writer stuck")
}
if s.Context().Err() == nil {
t.Fatal("writer failure retained native stream")
}
select {
case <-record.(*recordStream).readerCh:
case <-time.After(time.Second):
t.Fatal("writer failure retained reader")
}
}
func TestStreamConcurrentResetLocallyClosesBeforeNotification(t *testing.T) {
runtime := newStreamRuntime("reset-once")
var calls atomic.Int32
s := newStreamHandle(context.Background(), runtime, clientFileScope(), StreamOpenRequest{StreamID: "reset-once"}, 0, nil, nil, 0, nil,
func(ctx context.Context, stream *streamHandle, _ string) error {
calls.Add(1)
if stream.Context().Err() == nil {
t.Error("notification preceded local teardown")
}
if _, retained := runtime.lookup(clientFileScope(), stream.ID()); retained {
t.Error("runtime retained stream during notification")
}
<-ctx.Done()
return ctx.Err()
}, nil, defaultStreamConfig())
if err := runtime.register(clientFileScope(), s); err != nil {
t.Fatal(err)
}
if err := s.SetWriteDeadline(time.Now().Add(30 * time.Millisecond)); err != nil {
t.Fatal(err)
}
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(1)
go func() { defer wg.Done(); _ = s.Reset(io.ErrClosedPipe) }()
}
wg.Wait()
if calls.Load() != 1 {
t.Fatalf("reset notifications=%d", calls.Load())
}
}
func TestRecordFailureRejectsOversizedCode(t *testing.T) {
if _, err := encodeRecordErrorFrame(RecordFailure{FailedSeq: 1, Code: RecordErrorCode(bytes.Repeat([]byte("x"), 1<<16))}); !errors.Is(err, errRecordFrameInvalid) {
t.Fatalf("oversized failure code: %v", err)
}
}
func TestRegressionRecordFailRecordWriteFailureIsTerminal(t *testing.T) {
s := newStreamHandle(context.Background(), newStreamRuntime("review-fail"), clientFileScope(), StreamOpenRequest{StreamID: "review-fail"}, 0, nil, nil, 0, nil, nil,
func(context.Context, *streamHandle, []byte) error { return io.ErrUnexpectedEOF }, defaultStreamConfig())
t.Cleanup(func() { s.markReset(io.ErrClosedPipe) })
record, err := WrapStreamAsRecord(s, RecordOpenOptions{})
if err != nil {
t.Fatal(err)
}
defer record.Close()
payload, err := encodeRecordBatchFrame([]recordOutboundMessage{{Seq: 1, Payload: []byte("valid incoming record")}}, 0, false)
if err != nil {
t.Fatal(err)
}
if err := s.pushChunk(buildTransferFrame(payload)); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if _, err := record.ReadRecord(ctx); err != nil {
t.Fatal(err)
}
if err := record.FailRecord(1, RecordFailure{FailedSeq: 1, Code: RecordErrorCodeApplyFailed, Message: "disk full"}); !errors.Is(err, io.ErrUnexpectedEOF) {
t.Fatalf("failure send: %v", err)
}
seq, err := record.WriteRecord(context.Background(), []byte("must not admit more records"))
if err == nil {
t.Fatalf("failed record stream still accepts data: seq=%d", seq)
}
}