package notify import ( "context" "io" "time" ) const defaultRecordCloseTimeout = 5 * time.Second type recordCloseMode uint8 const ( recordCloseNone recordCloseMode = iota recordCloseWrite recordCloseFull ) func (r *recordStream) CloseWrite() error { if r == nil { return errRecordStreamNil } return r.closeRecord(recordCloseWrite) } func (r *recordStream) Close() error { if r == nil { return nil } return r.closeRecord(recordCloseFull) } func (r *recordStream) closeRecord(mode recordCloseMode) error { r.closeMu.Lock() defer r.closeMu.Unlock() if r.closeDone { return r.closeErr } if mode == recordCloseWrite && r.halfClosed { return nil } r.mu.Lock() r.outboundClosed = true target := r.enqueuedOutboundSeq err := r.streamErrorLocked() r.signalStateLocked() r.mu.Unlock() if err != nil { r.abortUnderlyingStream(err) r.closeDone, r.closeErr = true, err return err } timeout := r.cfg.CloseTimeout if timeout <= 0 { timeout = defaultRecordCloseTimeout } // A successful underlying Close can itself cancel the stream context. ctx, cancel := context.WithTimeout(context.Background(), timeout) defer cancel() err = r.flushAndClose(ctx, target, mode) if err != nil { r.setTerminalError(err) r.abortUnderlyingStream(err) r.closeDone, r.closeErr = true, err return err } if mode == recordCloseFull { r.setTerminalError(io.ErrClosedPipe) r.closeDone = true } else { r.halfClosed = true } return nil } func (r *recordStream) flushAndClose(ctx context.Context, target uint64, mode recordCloseMode) error { req := recordFlushRequest{ ctx: ctx, targetSeq: target, forceAck: true, closeMode: mode, done: make(chan error, 1), } select { case <-r.ctx.Done(): return r.streamError() case <-ctx.Done(): return ctx.Err() case r.flushCh <- req: } select { case err := <-req.done: return err case <-r.writerCh: select { case err := <-req.done: return err default: return r.streamError() } case <-ctx.Done(): select { case err := <-req.done: return err default: return ctx.Err() } } } // The writer serializes the last data, final ACK and close control message. func (r *recordStream) closeUnderlyingFromWriter(req recordFlushRequest) error { if err := req.ctx.Err(); err != nil { return err } if req.closeMode == recordCloseWrite && r.useHalfClose { return r.writePayloadFrame(encodeRecordFINFrame(req.targetSeq)) } deadline, _ := req.ctx.Deadline() if err := r.stream.SetWriteDeadline(deadline); err != nil { return err } if req.closeMode == recordCloseFull { return r.stream.Close() } return r.stream.CloseWrite() } func (r *recordStream) abortUnderlyingStream(err error) { if stream, ok := r.stream.(*streamHandle); ok { stream.mu.Lock() resetFn := stream.resetFn stream.mu.Unlock() // Local teardown must not wait for a reset reply from a stalled peer. if !stream.applyResetState(err) { return } if resetFn != nil { go func() { ctx, cancel := context.WithTimeout(context.Background(), defaultRecordCloseTimeout) defer cancel() _ = resetFn(ctx, stream, streamResetMessage(err)) }() } return } _ = r.stream.SetDeadline(time.Now()) _ = r.stream.Reset(err) } func (r *recordStream) abortRecord(err error) { r.setTerminalError(err) r.abortUnderlyingStream(err) } func (r *recordStream) abortProtocol(err error) { r.setTerminalError(err) _ = r.notifyFailureAndAbort(RecordFailure{ FailedSeq: r.nextInboundFailureSeq(), Code: RecordErrorCodeProtocol, Message: err.Error(), }) } func (r *recordStream) abortTimeout() time.Duration { if r.cfg.CloseTimeout > 0 && r.cfg.CloseTimeout < defaultRecordCloseTimeout { return r.cfg.CloseTimeout } return defaultRecordCloseTimeout } func (r *recordStream) closeReceive() { r.recvCloseOnce.Do(func() { close(r.recvCh) }) } func (r *recordStream) receiveFIN(finalSeq uint64) error { r.mu.Lock() if !r.useHalfClose || r.inboundClosed || finalSeq != r.inboundReceivedSeq { r.mu.Unlock() return errRecordSeqInvalid } r.inboundClosed = true r.mu.Unlock() // The reader owns recvCh; it keeps consuming control frames after data EOF. r.closeReceive() return nil }