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