From 1f2e74accae78b890afde56c93b1596cb6975b50 Mon Sep 17 00:00:00 2001 From: starainrt Date: Wed, 23 Sep 2026 15:33:17 +0800 Subject: [PATCH] =?UTF-8?q?fix(notify):=20=E4=BF=AE=E5=A4=8D=E4=BC=A0?= =?UTF-8?q?=E8=BE=93=E7=94=9F=E5=91=BD=E5=91=A8=E6=9C=9F=E7=AB=9E=E6=80=81?= =?UTF-8?q?=EF=BC=8C=E5=AE=8C=E5=96=84=E8=83=8C=E5=8E=8B=E4=B8=8E=E5=8D=8F?= =?UTF-8?q?=E8=AE=AE=E8=BE=B9=E7=95=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 完善 stream/bulk DataID 分配、预留和双向命名空间,修复并发打开及 dedicated/shared 回退时的 ID 冲突 - 将收发、回复、恢复任务和 sidecar 绑定原始会话与物理连接,防止重连后的旧消息误操作新连接 - 加强 close/reset 身份校验及实例移除检查,修复 dedicated attach 失败、通道引用和资源回收竞态 - 收紧批量发送器停止准入,确保在途入队完成后统一清理请求、缓冲区和等待者 - 修复 record 满队列死锁、取消时序号消耗及关闭竞态,确保关闭有界并返回真实错误 - 增加协商式 record 逻辑半关闭,保留反向 ACK;通过 reset 传递 RecordFailure,避免背压掩盖原始失败原因 - 补齐帧长度、批次数量、序号溢出和未确认窗口校验,提前拒绝超限数据并按字节预算拆批 - 为入站分发增加全局及单连接的条数、字节预算和阻塞背压,关闭时唤醒等待者,消除正常断连日志噪音 - 完善 bulk 窗口释放失败处理与传输诊断,补充并发、重连、背压、协议边界及真实 TCP 回归覆盖 --- batch_sender_error.go | 47 ++ bulk.go | 359 +++++++++++++-- bulk_attach_rejection_test.go | 39 ++ bulk_batch_sender.go | 137 ++++-- bulk_buffer_release_test.go | 68 +++ bulk_control.go | 210 +++++++-- bulk_dataid_test.go | 457 +++++++++++++++++++ bulk_dedicated.go | 451 +++++++++++++++++-- bulk_dedicated_attach_test.go | 245 ++++++++++ bulk_dedicated_batch.go | 247 +++++++--- bulk_dedicated_lane_sender.go | 162 +++++-- bulk_dedicated_sidecar.go | 283 ++++++++++-- bulk_dispatcher.go | 33 +- bulk_fastpath.go | 76 +++- bulk_lifecycle_lease_test.go | 183 ++++++++ bulk_recovery.go | 279 ++++++++++++ bulk_recovery_test.go | 719 ++++++++++++++++++++++++++++++ bulk_runtime.go | 459 +++++++++++++++++-- bulk_shared_batch.go | 52 ++- bulk_shared_batch_test.go | 32 ++ bulk_test.go | 139 ++++++ client.go | 3 + client_bulk.go | 240 +++++++--- client_config.go | 2 + client_conn.go | 3 +- client_conn_transport.go | 2 +- client_runtime.go | 34 +- client_send.go | 48 +- client_session_route.go | 108 +++++ client_session_runtime.go | 10 + client_session_runtime_test.go | 71 +++ client_stream.go | 96 ++-- client_transport.go | 82 +++- control_batch_sender.go | 21 +- file_receiver_test.go | 8 +- inbound_dispatcher.go | 177 ++++++-- inbound_dispatcher_test.go | 231 ++++++++++ logical_conn.go | 15 +- msg.go | 11 + protocol_bounds_test.go | 317 +++++++++++++ record_benchmark_test.go | 40 ++ record_codec.go | 71 ++- record_lifecycle.go | 184 ++++++++ record_negotiation.go | 28 +- record_network_test.go | 177 ++++++++ record_protocol_lifecycle_test.go | 318 +++++++++++++ record_regression_test.go | 373 ++++++++++++++++ record_reset.go | 23 + record_stream.go | 352 +++++++-------- record_stream_test.go | 94 +++- review_fix_regression_test.go | 197 ++++++++ server.go | 3 + server_bulk.go | 293 ++++++++---- server_inbound_source.go | 11 +- server_listen.go | 7 +- server_send.go | 10 +- server_session.go | 17 + server_stream.go | 110 ++--- session_owner_state.go | 67 ++- session_owner_state_test.go | 35 +- session_state.go | 8 +- signal_reliable.go | 20 +- stream.go | 163 +++++-- stream_batch_sender.go | 105 ++++- stream_control.go | 170 +++++-- stream_dispatcher.go | 50 ++- stream_fastpath.go | 24 +- stream_lifecycle_test.go | 64 +++ stream_route_dataid_test.go | 415 +++++++++++++++++ stream_runtime.go | 289 +++++++++++- stream_shared_batch.go | 48 +- transfer_observability_test.go | 8 +- transfer_plane.go | 6 + transfer_send_pipeline.go | 4 + transport_codec.go | 12 + transport_conn.go | 47 +- transport_conn_test.go | 43 ++ transport_write.go | 16 + transport_write_test.go | 145 ++++++ 79 files changed, 9190 insertions(+), 1013 deletions(-) create mode 100644 batch_sender_error.go create mode 100644 bulk_attach_rejection_test.go create mode 100644 bulk_dataid_test.go create mode 100644 bulk_lifecycle_lease_test.go create mode 100644 bulk_recovery.go create mode 100644 bulk_recovery_test.go create mode 100644 client_session_route.go create mode 100644 protocol_bounds_test.go create mode 100644 record_benchmark_test.go create mode 100644 record_lifecycle.go create mode 100644 record_network_test.go create mode 100644 record_protocol_lifecycle_test.go create mode 100644 record_regression_test.go create mode 100644 record_reset.go create mode 100644 review_fix_regression_test.go create mode 100644 stream_lifecycle_test.go create mode 100644 stream_route_dataid_test.go diff --git a/batch_sender_error.go b/batch_sender_error.go new file mode 100644 index 0000000..5f9e2b9 --- /dev/null +++ b/batch_sender_error.go @@ -0,0 +1,47 @@ +package notify + +import ( + "context" + "errors" +) + +// batchSenderQueueWaitError marks a request that expired or was canceled +// while waiting for the shared physical write gate. The transport is still +// healthy in this case, so the sender must not become permanently failed. +type batchSenderQueueWaitError struct { + err error +} + +func (e *batchSenderQueueWaitError) Error() string { + if e == nil || e.err == nil { + return "batch sender write queue wait failed" + } + return "batch sender write queue wait failed: " + e.err.Error() +} + +func (e *batchSenderQueueWaitError) Unwrap() error { + if e == nil { + return nil + } + return e.err +} + +func newBatchSenderQueueWaitError(err error) error { + if err == nil { + return nil + } + var existing *batchSenderQueueWaitError + if errors.As(err, &existing) { + return err + } + return &batchSenderQueueWaitError{err: err} +} + +func isBatchSenderQueueWaitError(err error) bool { + var queueErr *batchSenderQueueWaitError + return errors.As(err, &queueErr) +} + +func isBatchSenderQueueWaitCause(err error) bool { + return errors.Is(err, context.Canceled) || isTimeoutLikeError(err) +} diff --git a/bulk.go b/bulk.go index a26f0e6..0b42dae 100644 --- a/bulk.go +++ b/bulk.go @@ -3,6 +3,7 @@ package notify import ( "context" "errors" + "fmt" "io" "net" "strings" @@ -32,6 +33,8 @@ const ( defaultBulkAcceptReadyTimeout = 10 * time.Second defaultBulkResetNotifyTimeout = 30 * time.Second defaultBulkDataWriteTimeout = 2 * time.Minute + bulkWindowReleaseRetryDelay = 25 * time.Millisecond + bulkWindowReleaseShutdownGrace = 100 * time.Millisecond ) type BulkMetadata map[string]string @@ -163,6 +166,7 @@ var ( errBulkRejected = errors.New("bulk open rejected") errBulkReset = errors.New("bulk reset") errBulkDataIDEmpty = errors.New("bulk data id is empty") + errBulkDataIDExhausted = errors.New("bulk data id exhausted") errBulkDataPathNotReady = errors.New("bulk data path is not implemented yet") errBulkRangeInvalid = errors.New("bulk range is invalid") errBulkBackpressureExceeded = errors.New("bulk inbound backpressure exceeded") @@ -291,7 +295,10 @@ type bulkHandle struct { rangeSpec BulkRange metadata BulkMetadata sessionEpoch uint64 + clientRoute clientSessionRoute client *ClientCommon + debug atomic.Bool + debugSide string logical *LogicalConn transport *TransportConn transportGeneration uint64 @@ -316,8 +323,11 @@ type bulkHandle struct { writeCtxCancel context.CancelFunc createdAt time.Time - writeMu sync.Mutex - mu sync.Mutex + writeMu sync.Mutex + mu sync.Mutex + negotiationMu sync.RWMutex + finalizeOnce sync.Once + acceptState atomic.Uint32 // 0=pending, 1=dispatched/handled, 2=reset before dispatch writeQueue chan bulkAsyncWriteRequest writeWorkerDone chan struct{} @@ -355,6 +365,7 @@ type bulkHandle struct { dedicatedReady chan struct{} dedicatedWriteClosed bool dedicatedActiveLease bool + dedicatedLaneLease bool dedicatedState bulkDedicatedAttachState dedicatedAttempts uint32 dedicatedLastCode string @@ -411,6 +422,7 @@ func newBulkHandle(parent context.Context, runtime *bulkRuntime, runtimeScope st ctx: ctx, cancel: cancel, createdAt: time.Now(), + debugSide: bulkDebugSide(logical), readNotify: make(chan struct{}, 1), flowNotify: make(chan struct{}, 1), writeStateDone: make(chan struct{}), @@ -454,11 +466,20 @@ func (b *bulkHandle) fastPathVersionSnapshot() uint8 { if b == nil { return bulkFastPathVersionV1 } - b.mu.Lock() - defer b.mu.Unlock() + b.negotiationMu.RLock() + defer b.negotiationMu.RUnlock() return normalizeBulkFastPathVersion(b.fastPathVersion) } +func (b *bulkHandle) setFastPathVersion(version uint8) { + if b == nil { + return + } + b.negotiationMu.Lock() + b.fastPathVersion = normalizeBulkFastPathVersion(version) + b.negotiationMu.Unlock() +} + func (b *bulkHandle) FastPathVersion() uint8 { return b.fastPathVersionSnapshot() } @@ -502,9 +523,20 @@ func (b *bulkHandle) TransportGeneration() uint64 { if b == nil { return 0 } + b.negotiationMu.RLock() + defer b.negotiationMu.RUnlock() return b.transportGeneration } +func (b *bulkHandle) setTransportGeneration(generation uint64) { + if b == nil || generation == 0 { + return + } + b.negotiationMu.Lock() + b.transportGeneration = generation + b.negotiationMu.Unlock() +} + func (b *bulkHandle) Dedicated() bool { if b == nil { return false @@ -618,6 +650,9 @@ func (b *bulkHandle) installDedicatedSender(sender *bulkDedicatedSender) *bulkDe } b.dedicatedMu.Lock() defer b.dedicatedMu.Unlock() + if b.dedicatedState == bulkDedicatedAttachStateClosed { + return nil + } if b.dedicatedSender != nil { return b.dedicatedSender } @@ -689,6 +724,10 @@ func (b *bulkHandle) attachDedicatedConn(conn net.Conn) error { return net.ErrClosed } b.dedicatedMu.Lock() + if b.dedicatedState == bulkDedicatedAttachStateClosed { + b.dedicatedMu.Unlock() + return b.dedicatedAttachClosedError() + } if b.dedicatedConn != nil { b.dedicatedMu.Unlock() return errors.New("bulk dedicated conn already attached") @@ -718,6 +757,10 @@ func (b *bulkHandle) attachDedicatedConnShared(conn net.Conn) error { return net.ErrClosed } b.dedicatedMu.Lock() + if b.dedicatedState == bulkDedicatedAttachStateClosed { + b.dedicatedMu.Unlock() + return b.dedicatedAttachClosedError() + } if b.dedicatedConn != nil { if b.dedicatedConn == conn { b.dedicatedConnOwned = false @@ -754,6 +797,10 @@ func (b *bulkHandle) replaceDedicatedConn(conn net.Conn) (net.Conn, *bulkDedicat return nil, nil, net.ErrClosed } b.dedicatedMu.Lock() + if b.dedicatedState == bulkDedicatedAttachStateClosed { + b.dedicatedMu.Unlock() + return nil, nil, b.dedicatedAttachClosedError() + } oldConn := b.dedicatedConn oldOwned := b.dedicatedConnOwned oldSender := b.dedicatedSender @@ -786,6 +833,10 @@ func (b *bulkHandle) replaceDedicatedConnShared(conn net.Conn) (net.Conn, *bulkD return nil, nil, net.ErrClosed } b.dedicatedMu.Lock() + if b.dedicatedState == bulkDedicatedAttachStateClosed { + b.dedicatedMu.Unlock() + return nil, nil, b.dedicatedAttachClosedError() + } oldConn := b.dedicatedConn oldOwned := b.dedicatedConnOwned oldSender := b.dedicatedSender @@ -836,6 +887,13 @@ func (b *bulkHandle) bestEffortCloseDedicatedWriteHalf() { } } +func (b *bulkHandle) dedicatedAttachClosedError() error { + if err := b.resetErrSnapshot(); err != nil { + return err + } + return io.ErrClosedPipe +} + func (b *bulkHandle) dedicatedWriteHalfClosedSnapshot() bool { if b == nil { return false @@ -850,6 +908,16 @@ func (b *bulkHandle) setClientSnapshotOwner(client *ClientCommon) { return } b.client = client + if client != nil { + b.debug.Store(client.IsDebugMode()) + } +} + +func bulkDebugSide(logical *LogicalConn) string { + if logical != nil { + return "server" + } + return "client" } func (b *bulkHandle) clearDedicatedConn() (net.Conn, bool) { @@ -888,10 +956,35 @@ func (b *bulkHandle) releaseDedicatedActiveReserved() bool { return true } +func (b *bulkHandle) markDedicatedLaneReserved() { + if b == nil { + return + } + b.dedicatedMu.Lock() + b.dedicatedLaneLease = true + b.dedicatedMu.Unlock() +} + +func (b *bulkHandle) releaseDedicatedLaneReserved() bool { + if b == nil { + return false + } + b.dedicatedMu.Lock() + defer b.dedicatedMu.Unlock() + if !b.dedicatedLaneLease { + return false + } + b.dedicatedLaneLease = false + return true +} + func (b *bulkHandle) markAcceptDispatched() bool { if b == nil { return false } + if !b.acceptState.CompareAndSwap(0, 1) { + return false + } b.acceptMu.Lock() defer b.acceptMu.Unlock() if b.acceptDispatched { @@ -905,6 +998,7 @@ func (b *bulkHandle) markAcceptHandled() { if b == nil { return } + b.acceptState.CompareAndSwap(0, 1) b.acceptMu.Lock() b.acceptDispatched = true b.acceptMu.Unlock() @@ -1047,14 +1141,66 @@ func (b *bulkHandle) acceptsClientSessionEpoch(epoch uint64) bool { return b.sessionEpoch == epoch } +func (b *bulkHandle) setClientSessionRoute(route clientSessionRoute) { + if b == nil { + return + } + b.sessionEpoch = route.epoch + b.clientRoute = route +} + +func (b *bulkHandle) clientSessionRouteSnapshot() clientSessionRoute { + if b == nil { + return clientSessionRoute{} + } + if !b.clientRoute.bound() && b.client != nil { + return b.client.clientSessionRouteSnapshot() + } + return b.clientRoute +} + +func (b *bulkHandle) acceptsClientSessionRoute(route clientSessionRoute) bool { + if !b.acceptsClientSessionEpoch(route.epoch) { + return false + } + if b == nil || b.clientRoute.binding == nil || route.binding == nil { + return true + } + return b.clientRoute.binding == route.binding +} + func (b *bulkHandle) acceptsTransportGeneration(transport *TransportConn) bool { if b == nil { return false } - if b.transportGeneration == 0 || transport == nil { + generation := b.TransportGeneration() + if generation == 0 || transport == nil { return true } - return b.transportGeneration == transport.TransportGeneration() + return generation == transport.TransportGeneration() +} + +func (b *bulkHandle) acceptsCurrentTransport() bool { + if b == nil { + return false + } + if b.transport != nil { + return b.transport.IsCurrent() + } + if b.client != nil { + return b.client.clientSessionRouteCurrent(b.clientSessionRouteSnapshot()) + } + return true +} + +func (b *bulkHandle) acceptDispatchAllowed() bool { + if b == nil || b.acceptState.Load() == 2 { + return false + } + if err := b.resetErrSnapshot(); err != nil { + return false + } + return b.acceptsCurrentTransport() } func (b *bulkHandle) dataIDSnapshot() uint64 { @@ -1429,6 +1575,7 @@ func (b *bulkHandle) markReset(err error) { if b == nil { return } + b.acceptState.CompareAndSwap(0, 2) b.applyResetState(bulkResetError(err)) b.finalize() } @@ -1631,6 +1778,73 @@ func (b *bulkHandle) takePendingWindowRelease() (int64, int, bulkReleaseSender) return bytes, chunks, release } +func (b *bulkHandle) restorePendingWindowRelease(bytes int64, chunks int) { + if b == nil || (bytes <= 0 && chunks <= 0) { + return + } + b.mu.Lock() + b.pendingReleaseBytes += bytes + b.pendingReleaseChunks += chunks + b.mu.Unlock() +} + +func (b *bulkHandle) waitWindowReleaseRetry() bool { + if b == nil { + return false + } + timer := time.NewTimer(bulkWindowReleaseRetryDelay) + defer timer.Stop() + select { + case <-timer.C: + return true + case <-b.Context().Done(): + return false + } +} + +func (b *bulkHandle) shouldResetAfterWindowReleaseFailure() bool { + if b == nil { + return false + } + b.mu.Lock() + defer b.mu.Unlock() + return b.resetErr == nil && !b.remoteClosed && !b.peerReadClosed && !b.localReadClosed +} + +func (b *bulkHandle) windowReleaseClosing() bool { + if b == nil { + return false + } + b.mu.Lock() + defer b.mu.Unlock() + return b.localClosed || b.remoteClosed || b.peerReadClosed || b.localReadClosed +} + +func (b *bulkHandle) waitWindowReleaseShutdown() bool { + if b == nil || !b.windowReleaseClosing() { + return false + } + timer := time.NewTimer(bulkWindowReleaseShutdownGrace) + defer timer.Stop() + select { + case <-b.Context().Done(): + return true + case <-timer.C: + return !b.shouldResetAfterWindowReleaseFailure() + } +} + +func isBulkWindowReleaseClosedError(err error) bool { + if err == nil { + return false + } + if errors.Is(err, io.ErrClosedPipe) || errors.Is(err, net.ErrClosed) { + return true + } + message := strings.ToLower(err.Error()) + return strings.Contains(message, "closed pipe") || strings.Contains(message, "closed network connection") +} + func (b *bulkHandle) runWindowReleaseLoop() { if b == nil { return @@ -1647,12 +1861,61 @@ func (b *bulkHandle) runWindowReleaseLoop() { if release == nil || (bytes <= 0 && chunks <= 0) { break } - _ = release(b, bytes, chunks) + debug := b.debugEnabled() + var releaseStarted time.Time + if debug { + releaseStarted = time.Now() + b.debugf("release begin bytes=%d chunks=%d", bytes, chunks) + } + err := release(b, bytes, chunks) + if debug { + b.mu.Lock() + pendingBytes, pendingChunks := b.pendingReleaseBytes, b.pendingReleaseChunks + b.mu.Unlock() + b.debugf("release end bytes=%d chunks=%d elapsed=%s pending-bytes=%d pending-chunks=%d error=%v", bytes, chunks, time.Since(releaseStarted), pendingBytes, pendingChunks, err) + } + if err != nil { + b.restorePendingWindowRelease(bytes, chunks) + if b.Context().Err() != nil { + return + } + if errors.Is(err, context.Canceled) || isTimeoutLikeError(err) { + if !b.waitWindowReleaseRetry() { + return + } + b.scheduleWindowRelease() + continue + } + if isBulkWindowReleaseClosedError(err) && b.waitWindowReleaseShutdown() { + return + } + if b.shouldResetAfterWindowReleaseFailure() { + b.markReset(err) + } + return + } } } } -func (b *bulkHandle) acquireOutboundWindow(ctx context.Context, size int, chunks int) error { +func (b *bulkHandle) debugEnabled() bool { + if b == nil { + return false + } + if b.debug.Load() { + return true + } + if b.logical != nil && b.logical.server != nil { + return b.logical.server.IsDebugMode() + } + return false +} + +func (b *bulkHandle) debugf(format string, args ...interface{}) { + fmt.Printf("[bulk-debug] at=%s side=%s id=%s data=%d age=%s %s\n", time.Now().Format(time.RFC3339Nano), b.debugSide, b.id, b.dataID, time.Since(b.createdAt), fmt.Sprintf(format, args...)) +} + +func (b *bulkHandle) acquireOutboundWindow(ctx context.Context, size int, chunks int) (retErr error) { if b == nil || size <= 0 || !b.flowControlEnabled() { return nil } @@ -1663,6 +1926,8 @@ func (b *bulkHandle) acquireOutboundWindow(ctx context.Context, size int, chunks if chunks <= 0 { chunks = 1 } + debug := b.debugEnabled() + var waitStarted time.Time for { b.mu.Lock() if b.resetErr != nil { @@ -1696,7 +1961,13 @@ func (b *bulkHandle) acquireOutboundWindow(ctx context.Context, size int, chunks return nil } notify := b.flowNotify + avail, inFlight := b.outboundAvailBytes, b.outboundInFlight b.mu.Unlock() + if debug && waitStarted.IsZero() { + waitStarted = time.Now() + b.debugf("window wait begin need=%d chunks=%d avail=%d inflight=%d", size, chunks, avail, inFlight) + defer func() { b.debugf("window wait end elapsed=%s error=%v", time.Since(waitStarted), retErr) }() + } select { case <-notify: case <-ctx.Done(): @@ -1737,7 +2008,17 @@ func (b *bulkHandle) releaseOutboundWindow(bytes int64, chunks int) { if b == nil || !b.flowControlEnabled() { return } + debug := b.debugEnabled() + var lockStarted time.Time + if debug { + lockStarted = time.Now() + } b.mu.Lock() + var lockWait time.Duration + if debug { + lockWait = time.Since(lockStarted) + } + beforeBytes, beforeChunks := b.outboundAvailBytes, b.outboundInFlight if b.windowBytes > 0 && bytes > 0 { b.outboundAvailBytes += bytes maxAvail := int64(b.windowBytes) @@ -1752,7 +2033,11 @@ func (b *bulkHandle) releaseOutboundWindow(bytes int64, chunks int) { } } b.notifyFlowLocked() + afterBytes, afterChunks := b.outboundAvailBytes, b.outboundInFlight b.mu.Unlock() + if debug { + b.debugf("release received bytes=%d chunks=%d avail=%d->%d inflight=%d->%d lock-wait=%s", bytes, chunks, beforeBytes, afterBytes, beforeChunks, afterChunks, lockWait) + } } func (b *bulkHandle) bufferedChunkCountLocked() int { @@ -1790,7 +2075,7 @@ func (b *bulkHandle) snapshot() BulkSnapshot { snapshot := BulkSnapshot{ ID: b.id, DataID: b.dataID, - FastPathVersion: normalizeBulkFastPathVersion(b.fastPathVersion), + FastPathVersion: b.fastPathVersionSnapshot(), Scope: normalizeFileScope(b.runtimeScope), Range: b.rangeSpec, Metadata: cloneBulkMetadata(b.metadata), @@ -1804,7 +2089,7 @@ func (b *bulkHandle) snapshot() BulkSnapshot { DedicatedAttachLastCode: dedicatedLastCode, DedicatedDataStarted: dedicatedDataStarted, SessionEpoch: b.sessionEpoch, - TransportGeneration: b.transportGeneration, + TransportGeneration: b.TransportGeneration(), LocalClosed: b.localClosed, LocalReadClosed: b.localReadClosed, RemoteClosed: b.remoteClosed, @@ -1833,7 +2118,7 @@ func (b *bulkHandle) snapshot() BulkSnapshot { var diag snapshotBindingDiagnostics switch { case b.logical != nil || b.transport != nil: - diag = snapshotBindingDiagnosticsFromLogical(b.logical, b.transport, b.transportGeneration) + diag = snapshotBindingDiagnosticsFromLogical(b.logical, b.transport, b.TransportGeneration()) case b.client != nil: diag = snapshotBindingDiagnosticsFromClient(b.client, b.sessionEpoch) } @@ -1859,31 +2144,33 @@ func (b *bulkHandle) finalize() { if b == nil { return } - b.markDedicatedAttachClosed() - b.maybeSendWindowRelease(0, true) - if b.cancel != nil { - b.cancel() - } - if b.writeCtxCancel != nil { - b.writeCtxCancel() - } - sender := b.clearDedicatedSender() - conn, owned := b.clearDedicatedConn() - if conn != nil && owned { - _ = conn.Close() - } - if sender != nil { - sender.stop() - } - if b.client != nil && b.releaseDedicatedActiveReserved() { - b.client.releaseBulkDedicatedActiveSlot() - } - if b.client != nil { - b.client.releaseBulkDedicatedLane(b.dedicatedLaneIDSnapshot()) - } - if b.runtime != nil { - b.runtime.remove(b.runtimeScope, b.id) - } + b.finalizeOnce.Do(func() { + b.markDedicatedAttachClosed() + b.maybeSendWindowRelease(0, true) + if b.cancel != nil { + b.cancel() + } + if b.writeCtxCancel != nil { + b.writeCtxCancel() + } + sender := b.clearDedicatedSender() + conn, owned := b.clearDedicatedConn() + if conn != nil && owned { + _ = conn.Close() + } + if sender != nil { + sender.stop() + } + if b.client != nil && b.releaseDedicatedActiveReserved() { + b.client.releaseBulkDedicatedActiveSlot() + } + if b.client != nil && b.releaseDedicatedLaneReserved() { + b.client.releaseBulkDedicatedLaneAtRoute(b.dedicatedLaneIDSnapshot(), b.clientSessionRouteSnapshot()) + } + if b.runtime != nil { + b.runtime.remove(b.runtimeScope, b) + } + }) } func (b *bulkHandle) recordReadLocked(n int, now time.Time) { diff --git a/bulk_attach_rejection_test.go b/bulk_attach_rejection_test.go new file mode 100644 index 0000000..7f20656 --- /dev/null +++ b/bulk_attach_rejection_test.go @@ -0,0 +1,39 @@ +package notify + +import ( + "context" + "errors" + "net" + "testing" +) + +func TestDedicatedAttachPreservesRemoteRejection(t *testing.T) { + for _, op := range []string{"attach", "attach-shared", "replace", "replace-shared"} { + t.Run(op, func(t *testing.T) { + bulk := newBulkHandle(context.Background(), nil, clientFileScope(), BulkOpenRequest{BulkID: "rejected", DataID: 1, Dedicated: true}, 0, nil, nil, 0, nil, nil, nil, nil, nil) + rejection := errors.New("remote application rejected bulk") + bulk.markAcceptReady(rejection) + bulk.markReset(rejection) + left, right := net.Pipe() + defer left.Close() + defer right.Close() + var err error + switch op { + case "attach": + err = bulk.attachDedicatedConn(left) + case "attach-shared": + err = bulk.attachDedicatedConnShared(left) + case "replace": + _, _, err = bulk.replaceDedicatedConn(left) + case "replace-shared": + _, _, err = bulk.replaceDedicatedConnShared(left) + } + if !errors.Is(err, rejection) { + t.Fatalf("attach lost rejection: %v", err) + } + if bulk.dedicatedConnSnapshot() != nil { + t.Fatal("rejected bulk retained connection") + } + }) + } +} diff --git a/bulk_batch_sender.go b/bulk_batch_sender.go index 20f66f4..b8ceff5 100644 --- a/bulk_batch_sender.go +++ b/bulk_batch_sender.go @@ -61,11 +61,14 @@ type bulkBatchSender struct { stopCh chan struct{} doneCh chan struct{} - stopOnce sync.Once - flushMu sync.Mutex - queued atomic.Int64 - errMu sync.Mutex - err error + stopOnce sync.Once + admissionMu sync.Mutex + admitting sync.WaitGroup + admissionClosed bool + flushMu sync.Mutex + queued atomic.Int64 + errMu sync.Mutex + err error } func newBulkBatchSender(binding *transportBinding, codec bulkBatchCodec, writeTimeoutProvider func() time.Duration) *bulkBatchSender { @@ -193,20 +196,15 @@ func (s *bulkBatchSender) submitFramesOwned(ctx context.Context, frames []bulkFa } req = cloneQueuedBulkBatchRequest(req) s.queued.Add(1) - select { - case <-ctx.Done(): + if !s.enqueue(req) { s.queued.Add(-1) if req.release != nil { req.release() } - return normalizeStreamDeadlineError(ctx.Err()) - case <-s.stopCh: - s.queued.Add(-1) - if req.release != nil { - req.release() + if err := ctx.Err(); err != nil { + return normalizeStreamDeadlineError(err) } return s.stoppedErr() - case s.reqCh <- req: } select { case err := <-req.done: @@ -261,7 +259,11 @@ func (s *bulkBatchSender) tryDirectSubmit(req bulkBatchRequest) (bool, error) { } err := s.flush([]bulkBatchRequest{req}) if err != nil { - s.setErr(err) + if isBatchSenderQueueWaitError(err) { + return true, err + } + s.markFailed(err) + s.waitAdmissions() s.failPending(err) return true, err } @@ -288,7 +290,10 @@ func (s *bulkBatchSender) run() { if timerCh == nil { select { case <-s.stopCh: - s.failPending(s.stoppedErr()) + err := s.stoppedErr() + s.waitAdmissions() + s.failBatch(batch, err) + s.failPending(err) return case next := <-s.reqCh: batch = append(batch, next) @@ -303,7 +308,10 @@ func (s *bulkBatchSender) run() { if timer != nil { timer.Stop() } - s.failPending(s.stoppedErr()) + err := s.stoppedErr() + s.waitAdmissions() + s.failBatch(batch, err) + s.failPending(err) return case next := <-s.reqCh: batch = append(batch, next) @@ -344,10 +352,17 @@ func (s *bulkBatchSender) run() { } s.flushMu.Unlock() if err != nil { - s.setErr(err) + if isBatchSenderQueueWaitError(err) { + for _, item := range active { + s.finishRequest(item, err) + } + continue + } + s.markFailed(err) for _, item := range active { s.finishRequest(item, err) } + s.waitAdmissions() s.failPending(err) return } @@ -360,6 +375,7 @@ func (s *bulkBatchSender) run() { func (s *bulkBatchSender) nextRequest() (bulkBatchRequest, bool) { select { case <-s.stopCh: + s.waitAdmissions() s.failPending(s.stoppedErr()) return bulkBatchRequest{}, false case req := <-s.reqCh: @@ -433,6 +449,9 @@ func (s *bulkBatchSender) flush(requests []bulkBatchRequest) error { lockAcquired, err := s.binding.withConnWriteLockContextStopDeadlineManaged(context.Background(), s.stopCh, writeDeadline, func(conn net.Conn) error { return writeFramedPayloadBatchUnlocked(conn, queue, frames) }) + if !lockAcquired && isBatchSenderQueueWaitCause(err) { + return newBatchSenderQueueWaitError(err) + } s.binding.observeBulkAdaptivePayloadWrite(payloadBytes, time.Since(started), writeTimeout, err) if lockAcquired && err != nil { // A failed framed write may have emitted only part of a frame. @@ -454,6 +473,15 @@ func (s *bulkBatchSender) encodeRequests(requests []bulkBatchRequest) ([]bulkBat return nil, nil } payloads := make([]bulkBatchEncodedPayload, 0, len(requests)) + released := false + defer func() { + if released { + return + } + for index := range payloads { + payloads[index].done() + } + }() batch := make([]bulkFastFrame, 0, minInt(len(requests), bulkFastBatchMaxItems)) mixedBatchLimit := s.sharedMixedPayloadLimit() batchRequestIndex := -1 @@ -465,6 +493,9 @@ func (s *bulkBatchSender) encodeRequests(requests []bulkBatchRequest) ([]bulkBat } payload, release, err := s.encodeBatch(batch) if err != nil { + if release != nil { + release() + } return err } payloads = append(payloads, bulkBatchEncodedPayload{payload: payload, release: release}) @@ -479,12 +510,15 @@ func (s *bulkBatchSender) encodeRequests(requests []bulkBatchRequest) ([]bulkBat for _, frame := range req.frames { if !bulkFastPathSupportsSharedBatch(req.fastPathVersion) { if err := flushBatch(); err != nil { - return nil, err + return payloads, err } batchBytes = bulkFastBatchHeaderLen payload, release, err := s.encodeSingle(frame) if err != nil { - return nil, err + if release != nil { + release() + } + return payloads, err } payloads = append(payloads, bulkBatchEncodedPayload{payload: payload, release: release}) continue @@ -492,12 +526,15 @@ func (s *bulkBatchSender) encodeRequests(requests []bulkBatchRequest) ([]bulkBat frameLen := bulkFastBatchFrameLen(frame) if frameLen+bulkFastBatchHeaderLen > bulkFastBatchMaxPlainBytes { if err := flushBatch(); err != nil { - return nil, err + return payloads, err } batchBytes = bulkFastBatchHeaderLen payload, release, err := s.encodeSingle(frame) if err != nil { - return nil, err + if release != nil { + release() + } + return payloads, err } payloads = append(payloads, bulkBatchEncodedPayload{payload: payload, release: release}) continue @@ -512,7 +549,7 @@ func (s *bulkBatchSender) encodeRequests(requests []bulkBatchRequest) ([]bulkBat } if len(batch) > 0 && (len(batch) >= bulkFastBatchMaxItems || batchBytes+frameLen > batchLimit) { if err := flushBatch(); err != nil { - return nil, err + return payloads, err } batchBytes = bulkFastBatchHeaderLen nextMixed = false @@ -529,8 +566,9 @@ func (s *bulkBatchSender) encodeRequests(requests []bulkBatchRequest) ([]bulkBat } } if err := flushBatch(); err != nil { - return nil, err + return payloads, err } + released = true return payloads, nil } @@ -621,10 +659,8 @@ func (s *bulkBatchSender) stop() { if s == nil { return } - s.stopOnce.Do(func() { - s.setErr(errTransportDetached) - close(s.stopCh) - }) + s.markFailed(errTransportDetached) + s.waitAdmissions() <-s.doneCh // Direct submissions flush on the caller goroutine rather than run(). Wait // for that path too before declaring the binding safe to hand off. @@ -632,6 +668,28 @@ func (s *bulkBatchSender) stop() { s.flushMu.Unlock() } +func (s *bulkBatchSender) enqueue(req bulkBatchRequest) bool { + if s == nil { + return false + } + s.admissionMu.Lock() + if s.admissionClosed { + s.admissionMu.Unlock() + return false + } + s.admitting.Add(1) + s.admissionMu.Unlock() + defer s.admitting.Done() + select { + case <-req.ctx.Done(): + return false + case <-s.stopCh: + return false + case s.reqCh <- req: + return true + } +} + func (s *bulkBatchSender) failPending(err error) { for { select { @@ -643,6 +701,12 @@ func (s *bulkBatchSender) failPending(err error) { } } +func (s *bulkBatchSender) failBatch(batch []bulkBatchRequest, err error) { + for _, item := range batch { + s.finishRequest(item, err) + } +} + func (s *bulkBatchSender) finishRequest(req bulkBatchRequest, err error) { if s != nil { s.queued.Add(-1) @@ -664,6 +728,25 @@ func (s *bulkBatchSender) setErr(err error) { s.errMu.Unlock() } +func (s *bulkBatchSender) markFailed(err error) { + if s == nil { + return + } + s.setErr(err) + s.stopOnce.Do(func() { + s.admissionMu.Lock() + s.admissionClosed = true + close(s.stopCh) + s.admissionMu.Unlock() + }) +} + +func (s *bulkBatchSender) waitAdmissions() { + if s != nil { + s.admitting.Wait() + } +} + func (s *bulkBatchSender) errSnapshot() error { if s == nil { return errTransportDetached diff --git a/bulk_buffer_release_test.go b/bulk_buffer_release_test.go index 64d517c..18cd96c 100644 --- a/bulk_buffer_release_test.go +++ b/bulk_buffer_release_test.go @@ -6,7 +6,9 @@ import ( "errors" "math" "net" + "strings" "sync" + "sync/atomic" "testing" "time" ) @@ -133,6 +135,72 @@ func TestBulkReadDoesNotBlockOnAsyncWindowRelease(t *testing.T) { } } +func TestBulkWindowReleaseRetriesTimeoutWithoutLosingCredit(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + var calls atomic.Int32 + secondAttempt := make(chan struct{}) + bulk := newBulkHandle(ctx, newBulkRuntime("buffer-release-retry"), clientFileScope(), BulkOpenRequest{ + BulkID: "buffer-release-retry", + DataID: 1, + ChunkSize: 4, + WindowBytes: 4, + MaxInFlight: 1, + }, 0, nil, nil, 0, nil, nil, nil, nil, func(_ *bulkHandle, bytes int64, chunks int) error { + if bytes != 4 || chunks != 1 { + t.Fatalf("release = (%d,%d), want (4,1)", bytes, chunks) + } + if calls.Add(1) == 1 { + return context.DeadlineExceeded + } + close(secondAttempt) + return nil + }) + defer bulk.finalize() + + bulk.maybeSendWindowRelease(4, true) + select { + case <-secondAttempt: + case <-time.After(time.Second): + t.Fatalf("window release was not retried, calls=%d", calls.Load()) + } + bulk.mu.Lock() + pendingBytes, pendingChunks, resetErr := bulk.pendingReleaseBytes, bulk.pendingReleaseChunks, bulk.resetErr + bulk.mu.Unlock() + if pendingBytes != 0 || pendingChunks != 0 { + t.Fatalf("pending release after successful retry=(%d,%d), want zero", pendingBytes, pendingChunks) + } + if resetErr != nil { + t.Fatalf("transient release timeout reset bulk: %v", resetErr) + } +} + +func TestBulkWindowReleasePermanentErrorResetsBulk(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + bulk := newBulkHandle(ctx, newBulkRuntime("buffer-release-reset-error"), clientFileScope(), BulkOpenRequest{ + BulkID: "buffer-release-reset-error", + DataID: 1, + ChunkSize: 4, + WindowBytes: 4, + MaxInFlight: 1, + }, 0, nil, nil, 0, nil, nil, nil, nil, func(_ *bulkHandle, bytes int64, chunks int) error { + if bytes != 4 || chunks != 1 { + t.Fatalf("release = (%d,%d), want (4,1)", bytes, chunks) + } + return errors.New("permanent release failure") + }) + bulk.maybeSendWindowRelease(4, true) + select { + case <-bulk.releaseWorkerDone: + case <-time.After(time.Second): + t.Fatal("window release worker did not stop after permanent error") + } + if err := bulk.resetErrSnapshot(); err == nil || !strings.Contains(err.Error(), "permanent release failure") { + t.Fatalf("reset error=%v, want permanent release failure", err) + } +} + func TestLegacyBulkReleaseHonorsBulkCancellation(t *testing.T) { client := NewClient().(*ClientCommon) if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { diff --git a/bulk_control.go b/bulk_control.go index 3b366ed..3c5dc0c 100644 --- a/bulk_control.go +++ b/bulk_control.go @@ -3,6 +3,7 @@ package notify import ( "context" "errors" + "strings" "time" ) @@ -35,6 +36,7 @@ type BulkOpenResponse struct { type BulkCloseRequest struct { BulkID string + DataID uint64 Full bool } @@ -157,6 +159,14 @@ func dispatchBulkAccept(handler func(BulkAcceptInfo) error, bulk *bulkHandle, in if bulk == nil { return errBulkNotFound } + if !bulk.acceptDispatchAllowed() { + if resetErr := bulk.resetErrSnapshot(); resetErr != nil { + return resetErr + } + err := transportDetachedErrorForTransport(bulk.TransportConn()) + bulk.markReset(err) + return err + } if !bulk.markAcceptDispatched() { return nil } @@ -178,6 +188,7 @@ func dispatchBulkAccept(handler func(BulkAcceptInfo) error, bulk *bulkHandle, in } func (c *ClientCommon) clientBulkAcceptReadyNotifier(bulk *bulkHandle) func(error) { + route := bulk.clientSessionRouteSnapshot() return func(readyErr error) { if c == nil || bulk == nil { return @@ -191,7 +202,7 @@ func (c *ClientCommon) clientBulkAcceptReadyNotifier(bulk *bulkHandle) func(erro } ctx, cancel := context.WithTimeout(context.Background(), defaultBulkAcceptReadyTimeout) defer cancel() - if _, err := sendBulkReadyClient(ctx, c, req); err != nil && bulk.Context().Err() == nil { + if _, err := sendBulkReadyClientAtRoute(ctx, c, route, req); err != nil && bulk.Context().Err() == nil { bulk.markReset(err) } } @@ -202,11 +213,8 @@ func sendBulkReadyServer(ctx context.Context, s *ServerCommon, logical *LogicalC return errBulkServerNil } if transport != nil { - if _, err := sendBulkReadyServerTransport(ctx, s, transport, req); err == nil { - return nil - } else if !errors.Is(err, errTransportDetached) && !errors.Is(err, errBulkTransportNil) { - return err - } + _, err := sendBulkReadyServerTransport(ctx, s, transport, req) + return err } if logical == nil { return errBulkLogicalConnNil @@ -296,6 +304,15 @@ func (c *ClientCommon) handleInboundBulkOpen(msg *Message) { replyBulkControlIfNeeded(msg, resp) return } + route := msg.clientRoute + if !route.bound() { + route = c.clientSessionRouteSnapshot() + } + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + resp.Error = err.Error() + replyBulkControlIfNeeded(msg, resp) + return + } if req.Dedicated { if err := clientDedicatedBulkSupportError(c); err != nil { resp.Error = err.Error() @@ -310,17 +327,38 @@ func (c *ClientCommon) handleInboundBulkOpen(msg *Message) { return } scope := clientFileScope() - if req.DataID == 0 { - req.DataID = runtime.nextDataID() - resp.DataID = req.DataID + if existing, ok := runtime.lookup(scope, req.BulkID); ok && !existing.acceptsClientSessionRoute(route) { + existing.markReset(transportDetachedSessionEpochError()) } if req.Dedicated && req.AttachToken == "" { req.AttachToken = newBulkAttachToken() } resp.AttachToken = req.AttachToken - bulk := newBulkHandle(c.clientStopContextSnapshot(), runtime, scope, req, c.currentClientSessionEpoch(), nil, nil, 0, clientBulkCloseSender(c), clientBulkResetSender(c), clientBulkDataSender(c, c.currentClientSessionEpoch()), clientBulkWriteSender(c, c.currentClientSessionEpoch()), clientBulkReleaseSender(c)) + bulk := newBulkHandle(clientSessionRouteContext(route), runtime, scope, req, route.epoch, nil, nil, 0, clientBulkCloseSender(c), clientBulkResetSender(c), clientBulkDataSender(c, route), clientBulkWriteSender(c, route), clientBulkReleaseSender(c)) bulk.setClientSnapshotOwner(c) - if err := runtime.register(scope, bulk); err != nil { + bulk.setClientSessionRoute(route) + if req.Dedicated { + if err := c.retainBulkDedicatedLaneAtRoute(bulk.dedicatedLaneIDSnapshot(), route); err != nil { + // newBulkHandle starts its write/release workers before the lane + // retain can fail. Reset the unadopted candidate so those workers and + // its context are reclaimed; no lane lease was acquired yet. + bulk.markReset(err) + resp.Error = err.Error() + replyBulkControlIfNeeded(msg, resp) + return + } + bulk.markDedicatedLaneReserved() + } + if err := runtime.adoptInbound(scope, bulk); err != nil { + resp.Error = err.Error() + replyBulkControlIfNeeded(msg, resp) + return + } + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + // Reattach may have won the race after the preflight check. Remove the + // just-registered handle before any handler or sidecar side effect. + runtime.remove(scope, bulk) + bulk.markReset(err) resp.Error = err.Error() replyBulkControlIfNeeded(msg, resp) return @@ -376,6 +414,11 @@ func (s *ServerCommon) handleInboundBulkOpen(msg *Message) { return } transport := messageTransportConnSnapshot(msg) + if transport != nil && !transport.IsCurrent() { + resp.Error = transportDetachedErrorForTransport(transport).Error() + replyBulkControlIfNeeded(msg, resp) + return + } if req.Dedicated { if err := logicalDedicatedBulkSupportError(logical); err != nil { resp.Error = err.Error() @@ -391,20 +434,28 @@ func (s *ServerCommon) handleInboundBulkOpen(msg *Message) { } } scope := serverFileScope(logical) - if req.DataID == 0 { - req.DataID = runtime.nextDataID() - resp.DataID = req.DataID + if existing, ok := runtime.lookup(scope, req.BulkID); ok && !existing.acceptsTransportGeneration(transport) { + existing.markReset(transportDetachedGenerationMismatchError(existing.TransportGeneration(), transport)) } if req.Dedicated && req.AttachToken == "" { req.AttachToken = newBulkAttachToken() } resp.AttachToken = req.AttachToken bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, bulkTransportGeneration(logical, transport), serverBulkCloseSender(s, logical, transport), serverBulkResetSender(s, logical, transport), serverBulkDataSender(s, transport), serverBulkWriteSender(s, logical, transport), serverBulkReleaseSender(s, logical, transport)) - if err := runtime.register(scope, bulk); err != nil { + if err := runtime.adoptInbound(scope, bulk); err != nil { resp.Error = err.Error() replyBulkControlIfNeeded(msg, resp) return } + if transport != nil && !transport.IsCurrent() { + // Reattach may have won the race after the preflight check. Remove the + // just-registered handle before any handler or sidecar side effect. + runtime.remove(scope, bulk) + bulk.markReset(transportDetachedErrorForTransport(transport)) + resp.Error = transportDetachedErrorForTransport(transport).Error() + replyBulkControlIfNeeded(msg, resp) + return + } s.attachServerDedicatedSidecarIfExists(logical, bulk) if runtime.handlerSnapshot() == nil { bulk.markReset(errBulkHandlerNotConfigured) @@ -447,12 +498,17 @@ func (c *ClientCommon) handleInboundBulkClose(msg *Message) { replyBulkControlIfNeeded(msg, resp) return } - bulk, ok := runtime.lookup(clientFileScope(), req.BulkID) + bulk, ok := runtime.lookupControl(clientFileScope(), req.BulkID, req.DataID) if !ok { resp.Error = errBulkNotFound.Error() replyBulkControlIfNeeded(msg, resp) return } + if !bulk.acceptsClientSessionRoute(msg.clientRoute) { + resp.Error = transportDetachedSessionEpochError().Error() + replyBulkControlIfNeeded(msg, resp) + return + } if req.Full { bulk.markPeerClosed() } else { @@ -478,12 +534,17 @@ func (s *ServerCommon) handleInboundBulkClose(msg *Message) { } logical := messageLogicalConnSnapshot(msg) scope := serverFileScope(logical) - bulk, ok := runtime.lookup(scope, req.BulkID) + bulk, ok := runtime.lookupControl(scope, req.BulkID, req.DataID) if !ok { resp.Error = errBulkNotFound.Error() replyBulkControlIfNeeded(msg, resp) return } + if !bulk.acceptsTransportGeneration(messageTransportConnSnapshot(msg)) { + resp.Error = transportDetachedGenerationMismatchError(bulk.TransportGeneration(), messageTransportConnSnapshot(msg)).Error() + replyBulkControlIfNeeded(msg, resp) + return + } if req.Full { bulk.markPeerClosed() } else { @@ -507,15 +568,17 @@ func (c *ClientCommon) handleInboundBulkReset(msg *Message) { replyBulkControlIfNeeded(msg, resp) return } - bulk, ok := runtime.lookup(clientFileScope(), req.BulkID) - if !ok && req.DataID != 0 { - bulk, ok = runtime.lookupByDataID(clientFileScope(), req.DataID) - } + bulk, ok := runtime.lookupControl(clientFileScope(), req.BulkID, req.DataID) if !ok { resp.Error = errBulkNotFound.Error() replyBulkControlIfNeeded(msg, resp) return } + if !bulk.acceptsClientSessionRoute(msg.clientRoute) { + resp.Error = transportDetachedSessionEpochError().Error() + replyBulkControlIfNeeded(msg, resp) + return + } if resp.BulkID == "" { resp.BulkID = bulk.ID() } @@ -533,13 +596,13 @@ func (c *ClientCommon) handleInboundBulkRelease(msg *Message) { if runtime == nil { return } - bulk, ok := runtime.lookup(clientFileScope(), req.BulkID) - if !ok && req.DataID != 0 { - bulk, ok = runtime.lookupByDataID(clientFileScope(), req.DataID) - } + bulk, ok := runtime.lookupControl(clientFileScope(), req.BulkID, req.DataID) if !ok { return } + if !bulk.acceptsClientSessionRoute(msg.clientRoute) { + return + } bulk.releaseOutboundWindow(req.Bytes, req.Chunks) } @@ -557,15 +620,17 @@ func (c *ClientCommon) handleInboundBulkReady(msg *Message) { replyBulkControlIfNeeded(msg, resp) return } - bulk, ok := runtime.lookup(clientFileScope(), req.BulkID) - if !ok && req.DataID != 0 { - bulk, ok = runtime.lookupByDataID(clientFileScope(), req.DataID) - } + bulk, ok := runtime.lookupControl(clientFileScope(), req.BulkID, req.DataID) if !ok { resp.Error = errBulkNotFound.Error() replyBulkControlIfNeeded(msg, resp) return } + if !bulk.acceptsClientSessionRoute(msg.clientRoute) { + resp.Error = transportDetachedSessionEpochError().Error() + replyBulkControlIfNeeded(msg, resp) + return + } if resp.BulkID == "" { resp.BulkID = bulk.ID() } @@ -594,15 +659,18 @@ func (s *ServerCommon) handleInboundBulkReset(msg *Message) { } logical := messageLogicalConnSnapshot(msg) scope := serverFileScope(logical) - bulk, ok := runtime.lookup(scope, req.BulkID) - if !ok && req.DataID != 0 { - bulk, ok = runtime.lookupByDataID(scope, req.DataID) - } + bulk, ok := runtime.lookupControl(scope, req.BulkID, req.DataID) if !ok { resp.Error = errBulkNotFound.Error() replyBulkControlIfNeeded(msg, resp) return } + transport := messageTransportConnSnapshot(msg) + if !bulk.acceptsTransportGeneration(transport) { + resp.Error = transportDetachedGenerationMismatchError(bulk.TransportGeneration(), transport).Error() + replyBulkControlIfNeeded(msg, resp) + return + } if resp.BulkID == "" { resp.BulkID = bulk.ID() } @@ -627,15 +695,18 @@ func (s *ServerCommon) handleInboundBulkReady(msg *Message) { } logical := messageLogicalConnSnapshot(msg) scope := serverFileScope(logical) - bulk, ok := runtime.lookup(scope, req.BulkID) - if !ok && req.DataID != 0 { - bulk, ok = runtime.lookupByDataID(scope, req.DataID) - } + bulk, ok := runtime.lookupControl(scope, req.BulkID, req.DataID) if !ok { resp.Error = errBulkNotFound.Error() replyBulkControlIfNeeded(msg, resp) return } + transport := messageTransportConnSnapshot(msg) + if !bulk.acceptsTransportGeneration(transport) { + resp.Error = transportDetachedGenerationMismatchError(bulk.TransportGeneration(), transport).Error() + replyBulkControlIfNeeded(msg, resp) + return + } if resp.BulkID == "" { resp.BulkID = bulk.ID() } @@ -659,13 +730,13 @@ func (s *ServerCommon) handleInboundBulkRelease(msg *Message) { } logical := messageLogicalConnSnapshot(msg) scope := serverFileScope(logical) - bulk, ok := runtime.lookup(scope, req.BulkID) - if !ok && req.DataID != 0 { - bulk, ok = runtime.lookupByDataID(scope, req.DataID) - } + bulk, ok := runtime.lookupControl(scope, req.BulkID, req.DataID) if !ok { return } + if !bulk.acceptsTransportGeneration(messageTransportConnSnapshot(msg)) { + return + } bulk.releaseOutboundWindow(req.Bytes, req.Chunks) } @@ -687,6 +758,17 @@ func sendBulkOpenClient(ctx context.Context, c Client, req BulkOpenRequest) (Bul return decodeBulkOpenResponse(msg) } +func sendBulkOpenClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req BulkOpenRequest) (BulkOpenResponse, error) { + if c == nil { + return BulkOpenResponse{}, errBulkClientNil + } + msg, err := c.sendObjCtxAtRoute(ctx, route, BulkOpenSignalKey, req) + if err != nil { + return BulkOpenResponse{}, err + } + return decodeBulkOpenResponse(msg) +} + func sendBulkOpenServerLogical(ctx context.Context, s Server, logical *LogicalConn, req BulkOpenRequest) (BulkOpenResponse, error) { if s == nil { return BulkOpenResponse{}, errBulkServerNil @@ -726,6 +808,17 @@ func sendBulkCloseClient(ctx context.Context, c Client, req BulkCloseRequest) (B return decodeBulkCloseResponse(msg) } +func sendBulkCloseClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req BulkCloseRequest) (BulkCloseResponse, error) { + if c == nil { + return BulkCloseResponse{}, errBulkClientNil + } + msg, err := c.sendObjCtxAtRoute(ctx, route, BulkCloseSignalKey, req) + if err != nil { + return BulkCloseResponse{}, err + } + return decodeBulkCloseResponse(msg) +} + func sendBulkCloseServerLogical(ctx context.Context, s Server, logical *LogicalConn, req BulkCloseRequest) (BulkCloseResponse, error) { if s == nil { return BulkCloseResponse{}, errBulkServerNil @@ -765,6 +858,17 @@ func sendBulkResetClient(ctx context.Context, c Client, req BulkResetRequest) (B return decodeBulkResetResponse(msg) } +func sendBulkResetClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req BulkResetRequest) (BulkResetResponse, error) { + if c == nil { + return BulkResetResponse{}, errBulkClientNil + } + msg, err := c.sendObjCtxAtRoute(ctx, route, BulkResetSignalKey, req) + if err != nil { + return BulkResetResponse{}, err + } + return decodeBulkResetResponse(msg) +} + func sendBulkResetServerLogical(ctx context.Context, s Server, logical *LogicalConn, req BulkResetRequest) (BulkResetResponse, error) { if s == nil { return BulkResetResponse{}, errBulkServerNil @@ -794,6 +898,10 @@ func sendBulkResetServerTransport(ctx context.Context, s Server, transport *Tran } func sendBulkReleaseClient(ctx context.Context, c *ClientCommon, req BulkReleaseRequest) error { + return sendBulkReleaseClientAtRoute(ctx, c, c.clientSessionRouteSnapshot(), req) +} + +func sendBulkReleaseClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req BulkReleaseRequest) error { if c == nil { return errBulkClientNil } @@ -801,11 +909,11 @@ func sendBulkReleaseClient(ctx context.Context, c *ClientCommon, req BulkRelease if err != nil { return err } - _, err = c.sendWithContext(ctx, TransferMsg{ + _, err = c.sendWithContextTimeoutAtRoute(ctx, route, TransferMsg{ Key: BulkReleaseSignalKey, Value: data, Type: MSG_ASYNC, - }) + }, 0) return err } @@ -974,6 +1082,9 @@ func bulkControlResultError(op string, accepted bool, message string, callErr er } func bulkControlMessageError(message string) error { + if message == errTransportDetached.Error() || strings.HasPrefix(message, errTransportDetached.Error()+":") { + return errTransportDetached + } switch message { case errBulkNotFound.Error(): return errBulkNotFound @@ -993,6 +1104,8 @@ func bulkControlMessageError(message string) error { return errBulkRangeInvalid case errBulkDataIDEmpty.Error(): return errBulkDataIDEmpty + case errBulkDataIDExhausted.Error(): + return errBulkDataIDExhausted default: return errors.New(message) } @@ -1023,6 +1136,17 @@ func sendBulkReadyClient(ctx context.Context, c Client, req BulkReadyRequest) (B return decodeBulkReadyResponse(msg) } +func sendBulkReadyClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req BulkReadyRequest) (BulkReadyResponse, error) { + if c == nil { + return BulkReadyResponse{}, errBulkClientNil + } + msg, err := c.sendObjCtxAtRoute(ctx, route, BulkReadySignalKey, req) + if err != nil { + return BulkReadyResponse{}, err + } + return decodeBulkReadyResponse(msg) +} + func sendBulkReadyServerLogical(ctx context.Context, s Server, logical *LogicalConn, req BulkReadyRequest) (BulkReadyResponse, error) { if s == nil { return BulkReadyResponse{}, errBulkServerNil diff --git a/bulk_dataid_test.go b/bulk_dataid_test.go new file mode 100644 index 0000000..af507d1 --- /dev/null +++ b/bulk_dataid_test.go @@ -0,0 +1,457 @@ +package notify + +import ( + "context" + "errors" + "math" + "sync" + "testing" + "time" +) + +func TestBulkRuntimeSeparatesBidirectionalDataIDNamespaces(t *testing.T) { + clientRuntime := newBulkRuntime("cblk") + serverRuntime := newBulkRuntime("sblk") + clientID, err := clientRuntime.reserveDataID("peer", 0) + if err != nil { + t.Fatalf("reserve client data id: %v", err) + } + serverID, err := serverRuntime.reserveDataID("peer", 0) + if err != nil { + t.Fatalf("reserve server data id: %v", err) + } + if clientID == serverID || clientID%2 != 1 || serverID%2 != 0 { + t.Fatalf("client/server data ids = %d/%d, want disjoint odd/even namespaces", clientID, serverID) + } + + clientBulk := newBulkHandle(context.Background(), clientRuntime, "peer", BulkOpenRequest{BulkID: "client", DataID: clientID}, 0, nil, nil, 0, nil, nil, nil, nil, nil) + serverBulk := newBulkHandle(context.Background(), serverRuntime, "peer", BulkOpenRequest{BulkID: "server", DataID: serverID}, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := clientRuntime.registerReserved("peer", clientBulk); err != nil { + t.Fatalf("register client bulk: %v", err) + } + if err := serverRuntime.registerReserved("peer", serverBulk); err != nil { + t.Fatalf("register server bulk: %v", err) + } + if got, ok := clientRuntime.lookupInboundFrame("peer", clientID); !ok || got != clientBulk { + t.Fatalf("client local frame lookup = %p/%v, want client bulk", got, ok) + } + if got, ok := clientRuntime.lookupInboundFrame("peer", serverID); ok || got != nil { + t.Fatalf("client peer frame lookup = %p/%v, want missing inbound bulk", got, ok) + } +} + +func TestBulkRuntimeRoutesLegacyZeroDataIDInboundFrames(t *testing.T) { + for _, role := range []string{"cblk", "sblk"} { + t.Run(role, func(t *testing.T) { + runtime := newBulkRuntime(role) + bulk := newBulkHandle(context.Background(), runtime, "peer", BulkOpenRequest{ + BulkID: "legacy-inbound", + // A legacy initiator leaves DataID unset and uses the ID + // allocated by the receiver's open response. + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerInbound("peer", bulk); err != nil { + t.Fatalf("register legacy inbound bulk: %v", err) + } + if got := bulk.dataIDSnapshot(); got == 0 { + t.Fatal("legacy inbound registration allocated zero data id") + } + got, ok := runtime.lookupInboundFrame("peer", bulk.dataIDSnapshot()) + if !ok || got != bulk { + t.Fatalf("legacy inbound frame lookup = %p/%v, want %p/true", got, ok, bulk) + } + }) + } +} + +func TestBulkRuntimeAllocatorSkipsLegacyInboundCollision(t *testing.T) { + runtime := newBulkRuntime("cblk") + inbound := newBulkHandle(context.Background(), runtime, "peer", BulkOpenRequest{ + BulkID: "legacy-inbound", + DataID: 1, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerInbound("peer", inbound); err != nil { + t.Fatalf("register legacy inbound bulk: %v", err) + } + reserved, err := runtime.reserveDataID("peer", 0) + if err != nil { + t.Fatalf("reserve outbound data id: %v", err) + } + if reserved == inbound.dataIDSnapshot() { + t.Fatalf("outbound allocator reused legacy inbound data id %d", reserved) + } +} + +func TestBulkRuntimeInboundExplicitDataIDCannotExhaustOutboundAllocator(t *testing.T) { + tests := []struct { + role string + poisonID uint64 + wantID uint64 + }{ + {role: "cblk", poisonID: math.MaxUint64, wantID: 1}, + {role: "sblk", poisonID: math.MaxUint64 - 1, wantID: 2}, + } + for _, test := range tests { + t.Run(test.role, func(t *testing.T) { + runtime := newBulkRuntime(test.role) + inbound := newBulkHandle(context.Background(), runtime, "peer", BulkOpenRequest{ + BulkID: "peer-controlled", + DataID: test.poisonID, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerInbound("peer", inbound); err != nil { + t.Fatalf("register inbound bulk: %v", err) + } + got, err := runtime.reserveDataID("peer", 0) + if err != nil { + t.Fatalf("reserve outbound data id after peer-controlled id: %v", err) + } + if got != test.wantID { + t.Fatalf("reserved outbound data id = %d, want %d", got, test.wantID) + } + }) + } +} + +func TestBulkRuntimeStaleFinalizeDoesNotRemoveReplacement(t *testing.T) { + runtime := newBulkRuntime("cblk") + old := newBulkHandle(context.Background(), runtime, "peer", BulkOpenRequest{ + BulkID: "reused", + DataID: 1, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerInbound("peer", old); err != nil { + t.Fatalf("register old bulk: %v", err) + } + old.markReset(errors.New("old failed")) + + replacement := newBulkHandle(context.Background(), runtime, "peer", BulkOpenRequest{ + BulkID: "reused", + DataID: 3, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerInbound("peer", replacement); err != nil { + t.Fatalf("register replacement bulk: %v", err) + } + + old.markReset(errors.New("late duplicate reset")) + if got, ok := runtime.lookup("peer", "reused"); !ok || got != replacement { + t.Fatalf("replacement bulk after stale finalize = %p/%v, want %p/true", got, ok, replacement) + } +} + +func TestBulkRuntimeControlLookupRejectsMismatchedIdentity(t *testing.T) { + runtime := newBulkRuntime("control") + inbound := newBulkHandle(context.Background(), runtime, "peer", BulkOpenRequest{BulkID: "inbound", DataID: 7}, 0, nil, nil, 0, nil, nil, nil, nil, nil) + outbound := newBulkHandle(context.Background(), runtime, "peer", BulkOpenRequest{BulkID: "outbound", DataID: 7}, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerInbound("peer", inbound); err != nil { + t.Fatalf("register inbound bulk: %v", err) + } + if err := runtime.registerOutbound("peer", outbound); err != nil { + t.Fatalf("register outbound bulk: %v", err) + } + if got, ok := runtime.lookupControl("peer", "outbound", 7); !ok || got != outbound { + t.Fatalf("matching control lookup = %p/%v, want outbound", got, ok) + } + if got, ok := runtime.lookupControl("peer", "outbound", 0); !ok || got != outbound { + t.Fatalf("BulkID-only control lookup = %p/%v, want outbound", got, ok) + } + if got, ok := runtime.lookupControl("peer", "outbound", 8); ok || got != nil { + t.Fatalf("mismatched control lookup = %p/%v, want rejection", got, ok) + } + if got, ok := runtime.lookupControl("peer", "", 7); ok || got != nil { + t.Fatalf("ambiguous data-only control lookup = %p/%v, want rejection", got, ok) + } +} + +func TestBulkCloseControlRejectsMismatchedDataID(t *testing.T) { + t.Run("client", func(t *testing.T) { + client := NewClient().(*ClientCommon) + runtime := client.getBulkRuntime() + bulk := newBulkHandle(context.Background(), runtime, clientFileScope(), BulkOpenRequest{ + BulkID: "close-client", + DataID: 11, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.register(clientFileScope(), bulk); err != nil { + t.Fatalf("register client bulk: %v", err) + } + defer bulk.markReset(errors.New("test cleanup")) + payload, err := encode(BulkCloseRequest{BulkID: bulk.ID(), DataID: 13, Full: true}) + if err != nil { + t.Fatalf("encode client close: %v", err) + } + client.handleInboundBulkClose(&Message{ + NetType: NET_CLIENT, + ServerConn: client, + TransferMsg: TransferMsg{Key: BulkCloseSignalKey, Value: payload, Type: MSG_ASYNC}, + }) + bulk.mu.Lock() + defer bulk.mu.Unlock() + if bulk.remoteClosed || bulk.peerReadClosed || bulk.resetErr != nil { + t.Fatalf("mismatched client close mutated bulk: remote=%v peer=%v reset=%v", bulk.remoteClosed, bulk.peerReadClosed, bulk.resetErr) + } + }) + + t.Run("server", func(t *testing.T) { + server := NewServer().(*ServerCommon) + logical := server.bootstrapAcceptedLogical("close-server", nil, nil) + if logical == nil { + t.Fatal("bootstrap server logical connection failed") + } + runtime := server.getBulkRuntime() + scope := serverFileScope(logical) + bulk := newBulkHandle(context.Background(), runtime, scope, BulkOpenRequest{ + BulkID: "close-server", + DataID: 17, + }, 0, logical, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.register(scope, bulk); err != nil { + t.Fatalf("register server bulk: %v", err) + } + defer bulk.markReset(errors.New("test cleanup")) + payload, err := encode(BulkCloseRequest{BulkID: bulk.ID(), DataID: 19, Full: true}) + if err != nil { + t.Fatalf("encode server close: %v", err) + } + server.handleInboundBulkClose(&Message{ + NetType: NET_SERVER, + LogicalConn: logical, + TransportConn: nil, + TransferMsg: TransferMsg{Key: BulkCloseSignalKey, Value: payload, Type: MSG_ASYNC}, + }) + bulk.mu.Lock() + defer bulk.mu.Unlock() + if bulk.remoteClosed || bulk.peerReadClosed || bulk.resetErr != nil { + t.Fatalf("mismatched server close mutated bulk: remote=%v peer=%v reset=%v", bulk.remoteClosed, bulk.peerReadClosed, bulk.resetErr) + } + }) +} + +func TestBulkOpenDedicatedConcurrentlyFromBothPeers(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKServer failed: %v", err) + } + serverAccepted := make(chan BulkAcceptInfo, 1) + server.SetBulkHandler(func(info BulkAcceptInfo) error { + serverAccepted <- info + return nil + }) + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatalf("server Listen failed: %v", err) + } + defer func() { _ = server.Stop() }() + + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKClient failed: %v", err) + } + clientAccepted := make(chan BulkAcceptInfo, 1) + client.SetBulkHandler(func(info BulkAcceptInfo) error { + clientAccepted <- info + return nil + }) + if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { + t.Fatalf("client Connect failed: %v", err) + } + defer func() { _ = client.Stop() }() + + var logical *LogicalConn + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + peers := server.GetLogicalConnList() + if len(peers) > 0 { + logical = peers[0] + break + } + time.Sleep(time.Millisecond) + } + if logical == nil { + t.Fatal("timed out waiting for server logical connection") + } + + type openResult struct { + bulk Bulk + err error + } + clientResult := make(chan openResult, 1) + serverResult := make(chan openResult, 1) + go func() { + bulk, err := client.OpenDedicatedBulk(context.Background(), BulkOpenOptions{Range: BulkRange{Length: 1}}) + clientResult <- openResult{bulk: bulk, err: err} + }() + go func() { + bulk, err := server.OpenBulkLogical(context.Background(), logical, BulkOpenOptions{Range: BulkRange{Offset: 1, Length: 1}}) + serverResult <- openResult{bulk: bulk, err: err} + }() + clientOpen := <-clientResult + serverOpen := <-serverResult + if clientOpen.err != nil || serverOpen.err != nil { + t.Fatalf("concurrent dedicated opens failed: client=%v server=%v", clientOpen.err, serverOpen.err) + } + clientInbound := waitAcceptedBulk(t, serverAccepted, 2*time.Second) + serverInbound := waitAcceptedBulk(t, clientAccepted, 2*time.Second) + if clientOpen.bulk.(*bulkHandle).dataIDSnapshot()%2 != 1 || serverOpen.bulk.(*bulkHandle).dataIDSnapshot()%2 != 0 { + t.Fatalf("local data ids = %d/%d, want odd/even", clientOpen.bulk.(*bulkHandle).dataIDSnapshot(), serverOpen.bulk.(*bulkHandle).dataIDSnapshot()) + } + if clientInbound.Bulk.(*bulkHandle).dataIDSnapshot() != clientOpen.bulk.(*bulkHandle).dataIDSnapshot() { + t.Fatalf("client-open data id mismatch across peers") + } + if serverInbound.Bulk.(*bulkHandle).dataIDSnapshot() != serverOpen.bulk.(*bulkHandle).dataIDSnapshot() { + t.Fatalf("server-open data id mismatch across peers") + } + if _, err := clientOpen.bulk.Write([]byte("client")); err != nil { + t.Fatalf("client bulk write failed: %v", err) + } + readBulkExactly(t, clientInbound.Bulk, "client", 2*time.Second) + if _, err := serverOpen.bulk.Write([]byte("server")); err != nil { + t.Fatalf("server bulk write failed: %v", err) + } + readBulkExactly(t, serverInbound.Bulk, "server", 2*time.Second) + _ = clientOpen.bulk.Close() + _ = clientInbound.Bulk.Close() + _ = serverOpen.bulk.Close() + _ = serverInbound.Bulk.Close() +} + +func TestBulkRuntimeDataIDAllocatorObservesExplicitIDs(t *testing.T) { + runtime := newBulkRuntime("dataid") + scope := "peer" + + explicit := newBulkHandle(context.Background(), runtime, scope, BulkOpenRequest{ + BulkID: "explicit", + DataID: 41, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.register(scope, explicit); err != nil { + t.Fatalf("register explicit bulk: %v", err) + } + if got := explicit.dataIDSnapshot(); got != 41 { + t.Fatalf("explicit data id = %d, want 41", got) + } + + auto := newBulkHandle(context.Background(), runtime, scope, BulkOpenRequest{ + BulkID: "auto", + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.register(scope, auto); err != nil { + t.Fatalf("register auto bulk: %v", err) + } + if got := auto.dataIDSnapshot(); got != 42 { + t.Fatalf("auto data id = %d, want 42 after explicit 41", got) + } + + reserved, err := runtime.reserveDataID(scope, 0) + if err != nil { + t.Fatalf("reserve data id: %v", err) + } + if reserved != 43 { + t.Fatalf("reserved data id = %d, want 43", reserved) + } + autoWhileReserved := newBulkHandle(context.Background(), runtime, scope, BulkOpenRequest{ + BulkID: "auto-while-reserved", + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.register(scope, autoWhileReserved); err != nil { + t.Fatalf("register auto bulk while id is reserved: %v", err) + } + if got := autoWhileReserved.dataIDSnapshot(); got != 44 { + t.Fatalf("auto data id while 43 is reserved = %d, want 44", got) + } + reservedBulk := newBulkHandle(context.Background(), runtime, scope, BulkOpenRequest{ + BulkID: "reserved", + DataID: reserved, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerReserved(scope, reservedBulk); err != nil { + t.Fatalf("register reserved bulk: %v", err) + } + + if _, err := runtime.reserveDataID(scope, reserved); err == nil { + t.Fatal("re-reserving an active data id should fail") + } +} + +func TestBulkRuntimeDataIDAllocatorConcurrentReservationsAreUnique(t *testing.T) { + runtime := newBulkRuntime("dataid-concurrent") + const count = 128 + ids := make(chan uint64, count) + errs := make(chan error, count) + var wg sync.WaitGroup + for i := 0; i < count; i++ { + wg.Add(1) + go func() { + defer wg.Done() + id, err := runtime.reserveDataID("peer", 0) + if err != nil { + errs <- err + return + } + ids <- id + }() + } + wg.Wait() + close(ids) + close(errs) + for err := range errs { + t.Fatalf("reserve data id: %v", err) + } + + seen := make(map[uint64]struct{}, count) + for id := range ids { + if _, ok := seen[id]; ok { + t.Fatalf("duplicate reserved data id %d", id) + } + seen[id] = struct{}{} + } + if len(seen) != count { + t.Fatalf("reserved id count = %d, want %d", len(seen), count) + } +} + +func TestBulkMixedDedicatedThenSharedUsesFreshDataID(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKServer failed: %v", err) + } + acceptCh := make(chan BulkAcceptInfo, 2) + server.SetBulkHandler(func(info BulkAcceptInfo) error { + acceptCh <- info + return nil + }) + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatalf("server Listen failed: %v", err) + } + defer func() { _ = server.Stop() }() + + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKClient failed: %v", err) + } + if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { + t.Fatalf("client Connect failed: %v", err) + } + defer func() { _ = client.Stop() }() + + first, err := client.OpenDedicatedBulk(context.Background(), BulkOpenOptions{ + Range: BulkRange{Offset: 0, Length: 1}, + }) + if err != nil { + t.Fatalf("open dedicated bulk: %v", err) + } + firstAccepted := waitAcceptedBulk(t, acceptCh, 2*time.Second) + firstID := first.(*bulkHandle).dataIDSnapshot() + if firstID == 0 || firstAccepted.Bulk.(*bulkHandle).dataIDSnapshot() != firstID { + t.Fatalf("dedicated data id mismatch: client=%d server=%d", firstID, firstAccepted.Bulk.(*bulkHandle).dataIDSnapshot()) + } + _ = first.Close() + _ = firstAccepted.Bulk.Close() + + second, err := client.OpenSharedBulk(context.Background(), BulkOpenOptions{ + Range: BulkRange{Offset: 1, Length: 1}, + }) + if err != nil { + t.Fatalf("open shared bulk after dedicated: %v", err) + } + secondAccepted := waitAcceptedBulk(t, acceptCh, 2*time.Second) + secondID := second.(*bulkHandle).dataIDSnapshot() + if secondID <= firstID { + t.Fatalf("shared data id = %d, want greater than dedicated id %d", secondID, firstID) + } + if got := secondAccepted.Bulk.(*bulkHandle).dataIDSnapshot(); got != secondID { + t.Fatalf("shared data id mismatch: client=%d server=%d", secondID, got) + } + _ = second.Close() + _ = secondAccepted.Bulk.Close() +} diff --git a/bulk_dedicated.go b/bulk_dedicated.go index ac10163..6c028f5 100644 --- a/bulk_dedicated.go +++ b/bulk_dedicated.go @@ -20,6 +20,7 @@ const ( systemBulkAttachKey = "_notify_bulk_attach" bulkDedicatedRecordMagic = "NBR1" bulkDedicatedRecordHeaderLen = 8 + bulkDedicatedRecordMaxBytes = 20 * 1024 * 1024 defaultBulkDedicatedAttachLimit = 16 defaultBulkDedicatedActiveLimit = 4096 @@ -187,6 +188,9 @@ func encodeDirectSignalFrame(queue *stario.StarQueue, sequenceEn func(interface{ if payload == nil && len(plain) != 0 { return nil, errTransportPayloadEncryptFailed } + if err := validateTransportFramePayloadLen(payload); err != nil { + return nil, err + } return queue.BuildMessage(payload), nil } @@ -210,7 +214,7 @@ func readDirectSignalFramePayload(conn net.Conn) ([]byte, error) { if conn == nil { return nil, net.ErrClosed } - return newTransportFrameReader(conn, stario.NewQueue()).Next() + return newTransportFrameReader(conn, stario.NewQueueCtx(nil, 1, transportFrameMaxPayloadBytes)).Next() } func writeBulkDedicatedRecord(conn net.Conn, payload []byte) error { @@ -218,17 +222,37 @@ func writeBulkDedicatedRecord(conn net.Conn, payload []byte) error { } func writeBulkDedicatedRecordWithDeadline(conn net.Conn, payload []byte, deadline time.Time) error { + return writeBulkDedicatedRecordWithDeadlineTrace(conn, payload, deadline, 0) +} + +func writeBulkDedicatedRecordWithDeadlineTrace(conn net.Conn, payload []byte, deadline time.Time, traceDataID uint64) error { if conn == nil { return net.ErrClosed } + if len(payload) > bulkDedicatedRecordMaxBytes { + return fmt.Errorf("%w: dedicated record payload=%d max=%d", errBulkFastPayloadInvalid, len(payload), bulkDedicatedRecordMaxBytes) + } if deadline.IsZero() { deadline = writeDeadlineFromTimeout(defaultBulkDataWriteTimeout) } + var prepareStarted time.Time + if traceDataID != 0 { + prepareStarted = time.Now() + fmt.Printf("[bulk-debug] at=%s record write prepare data=%d conn=%T local=%v remote=%v bytes=%d\n", prepareStarted.Format(time.RFC3339Nano), traceDataID, conn, conn.LocalAddr(), conn.RemoteAddr(), len(payload)) + prepareStarted = time.Now() + } return withRawConnWriteLockDeadline(conn, deadline, func(conn net.Conn) error { var header [bulkDedicatedRecordHeaderLen]byte copy(header[:4], bulkDedicatedRecordMagic) binary.BigEndian.PutUint32(header[4:8], uint32(len(payload))) - return writeNetBuffersFullUnlocked(conn, net.Buffers{header[:], payload}) + if traceDataID == 0 { + return writeNetBuffersFullUnlocked(conn, net.Buffers{header[:], payload}) + } + fmt.Printf("[bulk-debug] at=%s record socket write begin data=%d gate-and-deadline=%s\n", time.Now().Format(time.RFC3339Nano), traceDataID, time.Since(prepareStarted)) + writeStarted := time.Now() + err := writeNetBuffersFullUnlocked(conn, net.Buffers{header[:], payload}) + fmt.Printf("[bulk-debug] at=%s record socket write end data=%d elapsed=%s error=%v\n", time.Now().Format(time.RFC3339Nano), traceDataID, time.Since(writeStarted), err) + return err }) } @@ -244,22 +268,49 @@ func readBulkDedicatedRecord(conn net.Conn) ([]byte, error) { } func readBulkDedicatedRecordPooled(conn net.Conn) ([]byte, func(), error) { + return readBulkDedicatedRecordPooledTrace(conn, false) +} + +func readBulkDedicatedRecordPooledTrace(conn net.Conn, debug bool) ([]byte, func(), error) { if conn == nil { return nil, nil, net.ErrClosed } + var headerStarted time.Time + if debug { + headerStarted = time.Now() + fmt.Printf("[bulk-debug] at=%s record read header begin local=%v remote=%v\n", headerStarted.Format(time.RFC3339Nano), conn.LocalAddr(), conn.RemoteAddr()) + headerStarted = time.Now() + } var header [bulkDedicatedRecordHeaderLen]byte if _, err := io.ReadFull(conn, header[:]); err != nil { + if debug { + fmt.Printf("[bulk-debug] at=%s record read header failed elapsed=%s error=%v\n", time.Now().Format(time.RFC3339Nano), time.Since(headerStarted), err) + } return nil, nil, err } if string(header[:4]) != bulkDedicatedRecordMagic { return nil, nil, fmt.Errorf("%w: record magic=%x", errBulkFastPayloadInvalid, header[:4]) } - size := int(binary.BigEndian.Uint32(header[4:8])) - if size < 0 { + wireSize := binary.BigEndian.Uint32(header[4:8]) + if wireSize > bulkDedicatedRecordMaxBytes { return nil, nil, errBulkFastPayloadInvalid } + size := int(wireSize) + if debug { + fmt.Printf("[bulk-debug] at=%s record read payload begin bytes=%d header=%s\n", time.Now().Format(time.RFC3339Nano), size, time.Since(headerStarted)) + } payload := getModernPSKPayloadBuffer(size) - if _, err := io.ReadFull(conn, payload); err != nil { + var readStarted time.Time + if debug { + readStarted = time.Now() + fmt.Printf("[bulk-debug] at=%s record read buffer ready bytes=%d\n", readStarted.Format(time.RFC3339Nano), size) + readStarted = time.Now() + } + n, err := io.ReadFull(conn, payload) + if debug { + fmt.Printf("[bulk-debug] at=%s record read payload end bytes=%d/%d elapsed=%s error=%v\n", time.Now().Format(time.RFC3339Nano), n, size, time.Since(readStarted), err) + } + if err != nil { putModernPSKPayloadBuffer(payload) return nil, nil, err } @@ -302,37 +353,76 @@ func (c *ClientCommon) attachDedicatedBulkSidecar(ctx context.Context, bulk *bul if ctx == nil { ctx = context.Background() } - laneID := bulk.dedicatedLaneIDSnapshot() - releaseActiveSlot, err := c.acquireBulkDedicatedActiveSlot(ctx) - if err != nil { + route := bulk.clientSessionRouteSnapshot() + if !route.bound() { + route = c.clientSessionRouteSnapshot() + } + attachCtx, cancelAttach := context.WithCancel(ctx) + defer cancelAttach() + stopRoute := func() bool { return false } + if route.transportStopCtx != nil { + stopRoute = context.AfterFunc(route.transportStopCtx, cancelAttach) + defer stopRoute() + } + checkRoute := func() error { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return err + } + return nil + } + routeWaitError := func(waitErr error) error { + if routeErr := checkRoute(); routeErr != nil { + return routeErr + } + return waitErr + } + if err := checkRoute(); err != nil { return err } + laneID := bulk.dedicatedLaneIDSnapshot() + releaseActiveSlot, err := c.acquireBulkDedicatedActiveSlot(attachCtx) + if err != nil { + return routeWaitError(err) + } needReleaseActive := true defer func() { if needReleaseActive { releaseActiveSlot() } }() - if sidecar := c.clientDedicatedSidecarSnapshotForLane(laneID); sidecar != nil && sidecar.conn != nil { - if err := bulk.attachDedicatedConnShared(sidecar.conn); err == nil { + if sidecar := c.clientDedicatedSidecarSnapshotForLaneAtRoute(laneID, route); sidecar != nil { + if err := checkRoute(); err != nil { + return err + } + if err := sidecar.withConn(func(conn net.Conn) error { + return bulk.attachDedicatedConnShared(conn) + }); err == nil { bulk.markDedicatedActiveReserved() needReleaseActive = false return nil } } - _, flight, leader := c.beginClientDedicatedSidecarAttach(laneID) + _, flight, leader, beginErr := c.beginClientDedicatedSidecarAttach(laneID, route) + if beginErr != nil { + return beginErr + } if !leader { if flight == nil { return errTransportDetached } - if err := flight.wait(ctx); err != nil { + if err := flight.wait(attachCtx); err != nil { + return routeWaitError(err) + } + if err := checkRoute(); err != nil { return err } - sidecar := c.clientDedicatedSidecarSnapshotForLane(laneID) - if sidecar == nil || sidecar.conn == nil { + sidecar := c.clientDedicatedSidecarSnapshotForLaneAtRoute(laneID, route) + if sidecar == nil { return errTransportDetached } - if err := bulk.attachDedicatedConnShared(sidecar.conn); err != nil { + if err := sidecar.withConn(func(conn net.Conn) error { + return bulk.attachDedicatedConnShared(conn) + }); err != nil { return err } bulk.markDedicatedActiveReserved() @@ -346,14 +436,20 @@ func (c *ClientCommon) attachDedicatedBulkSidecar(ctx context.Context, bulk *bul defer func() { c.finishClientDedicatedSidecarAttach(laneID, flight, flightErr) }() - releaseAttachSlot, err := c.acquireBulkDedicatedAttachSlot(ctx) + releaseAttachSlot, err := c.acquireBulkDedicatedAttachSlot(attachCtx) if err != nil { - flightErr = err - return err + flightErr = routeWaitError(err) + return flightErr } defer releaseAttachSlot() - if sidecar := c.clientDedicatedSidecarSnapshotForLane(laneID); sidecar != nil && sidecar.conn != nil { - if err := bulk.attachDedicatedConnShared(sidecar.conn); err == nil { + if sidecar := c.clientDedicatedSidecarSnapshotForLaneAtRoute(laneID, route); sidecar != nil { + if err := checkRoute(); err != nil { + flightErr = err + return err + } + if err := sidecar.withConn(func(conn net.Conn) error { + return bulk.attachDedicatedConnShared(conn) + }); err == nil { bulk.markDedicatedActiveReserved() needReleaseActive = false flightErr = nil @@ -367,25 +463,36 @@ func (c *ClientCommon) attachDedicatedBulkSidecar(ctx context.Context, bulk *bul } var lastErr error for attempt := 1; attempt <= attempts; attempt++ { + if err := checkRoute(); err != nil { + flightErr = err + return err + } c.bulkAttachAttemptCount.Add(1) if attempt > 1 { delay := backoff * time.Duration(1<<(attempt-2)) if delay > 3*time.Second { delay = 3 * time.Second } - if err := waitDedicatedAttachBackoff(ctx, delay); err != nil { - flightErr = err - return err + if err := waitDedicatedAttachBackoff(attachCtx, delay); err != nil { + flightErr = routeWaitError(err) + return flightErr } } - dialCtx := ctx + dialCtx := attachCtx dialCancel := func() {} if dialTimeout > 0 { - dialCtx, dialCancel = context.WithTimeout(ctx, dialTimeout) + dialCtx, dialCancel = context.WithTimeout(attachCtx, dialTimeout) } conn, err := c.dialDedicatedBulkConn(dialCtx, dialTimeout) dialCancel() + if err == nil && conn == nil { + err = errTransportDetached + } if err != nil { + if routeErr := checkRoute(); routeErr != nil { + flightErr = routeErr + return routeErr + } lastErr = err if attempt < attempts && isRetryableDedicatedAttachError(err) { flightErr = err @@ -396,15 +503,24 @@ func (c *ClientCommon) attachDedicatedBulkSidecar(ctx context.Context, bulk *bul flightErr = err return err } - helloCtx := ctx + if err := checkRoute(); err != nil { + _ = conn.Close() + flightErr = err + return err + } + helloCtx := attachCtx helloCancel := func() {} if helloTimeout > 0 { - helloCtx, helloCancel = context.WithTimeout(ctx, helloTimeout) + helloCtx, helloCancel = context.WithTimeout(attachCtx, helloTimeout) } resp, err := c.sendDedicatedBulkAttachRequest(helloCtx, conn, bulk) helloCancel() if err != nil { _ = conn.Close() + if routeErr := checkRoute(); routeErr != nil { + flightErr = routeErr + return routeErr + } lastErr = err if attempt < attempts && isRetryableDedicatedAttachError(err) { flightErr = err @@ -416,6 +532,11 @@ func (c *ClientCommon) attachDedicatedBulkSidecar(ctx context.Context, bulk *bul flightErr = err return err } + if err := checkRoute(); err != nil { + _ = conn.Close() + flightErr = err + return err + } if !resp.Accepted { _ = conn.Close() rejectedErr := &bulkAttachError{ @@ -439,18 +560,30 @@ func (c *ClientCommon) attachDedicatedBulkSidecar(ctx context.Context, bulk *bul flightErr = rejectedErr return rejectedErr } + if err := checkRoute(); err != nil { + _ = conn.Close() + flightErr = err + return err + } sidecar := newBulkDedicatedSidecar(conn, laneID) - activeSidecar, installed := c.installClientDedicatedSidecar(laneID, sidecar) + activeSidecar, installed, installErr := c.installClientDedicatedSidecarAtRoute(laneID, sidecar, route) + if installErr != nil { + sidecar.close() + flightErr = installErr + return installErr + } if !installed { sidecar.close() sidecar = activeSidecar } - if sidecar == nil || sidecar.conn == nil { + if sidecar == nil { bulk.setDedicatedAttachLastCode(string(bulkAttachErrorCodeAttachFailed)) flightErr = errTransportDetached return errTransportDetached } - if err := bulk.attachDedicatedConnShared(sidecar.conn); err != nil { + if err := sidecar.withConn(func(conn net.Conn) error { + return bulk.attachDedicatedConnShared(conn) + }); err != nil { if installed && c.clearClientDedicatedSidecar(laneID, sidecar) { sidecar.close() } @@ -462,9 +595,17 @@ func (c *ClientCommon) attachDedicatedBulkSidecar(ctx context.Context, bulk *bul } return err } + if err := checkRoute(); err != nil { + if installed && c.clearClientDedicatedSidecar(laneID, sidecar) { + sidecar.close() + } + bulk.markReset(err) + flightErr = err + return err + } c.bulkAttachSuccessCount.Add(1) if installed { - go c.readDedicatedSidecarLoop(sidecar) + go c.readDedicatedSidecarLoopAtRoute(sidecar, route) } bulk.markDedicatedActiveReserved() needReleaseActive = false @@ -651,6 +792,9 @@ func (c *ClientCommon) sendDedicatedBulkAttachRequest(ctx context.Context, conn if bulk == nil { return bulkAttachResponse{}, errBulkIDEmpty } + if conn == nil { + return bulkAttachResponse{}, errTransportDetached + } if ctx == nil { ctx = context.Background() } @@ -716,15 +860,40 @@ func (c *ClientCommon) clientDedicatedBulkAttachTransportProtectionProfile() tra } func (c *ClientCommon) readDedicatedSidecarLoop(sidecar *bulkDedicatedSidecar) { + c.readDedicatedSidecarLoopAtRoute(sidecar, c.clientSessionRouteSnapshot()) +} + +func (c *ClientCommon) readDedicatedSidecarLoopAtRoute(sidecar *bulkDedicatedSidecar, route clientSessionRoute) { if c == nil || sidecar == nil || sidecar.conn == nil { return } + debug := c.IsDebugMode() + routeCurrent := func() bool { + if !c.clientSessionRouteCurrent(route) { + return false + } + return route.transportStopCtx == nil || route.transportStopCtx.Err() == nil + } for { - payload, payloadRelease, err := readBulkDedicatedRecordPooled(sidecar.conn) + if !routeCurrent() { + return + } + payload, payloadRelease, err := readBulkDedicatedRecordPooledTrace(sidecar.conn, debug) if err != nil { c.handleClientDedicatedSidecarFailure(sidecar, err) return } + if !routeCurrent() { + if payloadRelease != nil { + payloadRelease() + } + return + } + var dispatchStarted time.Time + if debug { + dispatchStarted = time.Now() + fmt.Printf("[bulk-debug] at=%s side=client lane=%d sidecar decode-dispatch begin bytes=%d\n", dispatchStarted.Format(time.RFC3339Nano), sidecar.laneID, len(payload)) + } profile := c.clientTransportProtectionSnapshot() plain, plainRelease, err := decryptTransportPayloadCodecPooled(profile.mode, profile.runtime, profile.msgDe, profile.secretKey, payload, payloadRelease) if err != nil { @@ -732,6 +901,10 @@ func (c *ClientCommon) readDedicatedSidecarLoop(sidecar *bulkDedicatedSidecar) { return } owner := newBulkReadPayloadOwner(plainRelease) + if !routeCurrent() { + owner.done() + return + } runtime := c.getBulkRuntime() if runtime == nil { owner.done() @@ -741,22 +914,29 @@ func (c *ClientCommon) readDedicatedSidecarLoop(sidecar *bulkDedicatedSidecar) { currentDataID uint64 currentBulk *bulkHandle skipDataID bool + staleRoute bool ) err = walkDedicatedBulkInboundPayload(plain, func(dataID uint64, item bulkDedicatedBatchItem) error { + if !routeCurrent() { + staleRoute = true + currentBulk = nil + skipDataID = true + return nil + } if dataID != currentDataID { currentDataID = dataID currentBulk = nil skipDataID = false - bulk, ok := runtime.lookupByDataID(clientFileScope(), dataID) + bulk, ok := runtime.lookupInboundFrame(clientFileScope(), dataID) if !ok { - c.bestEffortRejectInboundBulkData("", dataID, errBulkNotFound.Error()) + c.bestEffortRejectInboundBulkDataAtRoute(route, "", dataID, errBulkNotFound.Error()) skipDataID = true return nil } - if !bulk.acceptsClientSessionEpoch(c.currentClientSessionEpoch()) { + if !bulk.acceptsClientSessionRoute(route) { detachErr := transportDetachedSessionEpochError() bulk.markReset(detachErr) - c.bestEffortRejectInboundBulkData(bulk.ID(), dataID, detachErr.Error()) + c.bestEffortRejectInboundBulkDataAtRoute(route, bulk.ID(), dataID, detachErr.Error()) skipDataID = true return nil } @@ -766,6 +946,20 @@ func (c *ClientCommon) readDedicatedSidecarLoop(sidecar *bulkDedicatedSidecar) { if skipDataID || currentBulk == nil { return nil } + if !routeCurrent() { + staleRoute = true + currentBulk = nil + skipDataID = true + return nil + } + if !currentBulk.acceptsClientSessionRoute(route) { + detachErr := transportDetachedSessionEpochError() + currentBulk.markReset(detachErr) + c.bestEffortRejectInboundBulkDataAtRoute(route, currentBulk.ID(), dataID, detachErr.Error()) + currentBulk = nil + skipDataID = true + return nil + } var release func() if item.Type == bulkFastPayloadTypeData { release = owner.retainChunk() @@ -794,6 +988,12 @@ func (c *ClientCommon) readDedicatedSidecarLoop(sidecar *bulkDedicatedSidecar) { return } owner.done() + if debug { + fmt.Printf("[bulk-debug] at=%s side=client lane=%d sidecar decode-dispatch end elapsed=%s\n", time.Now().Format(time.RFC3339Nano), sidecar.laneID, time.Since(dispatchStarted)) + } + if staleRoute { + return + } } } @@ -802,6 +1002,16 @@ func (s *ServerCommon) handleBulkAttachSystemMessage(message Message) bool { return false } current := messageLogicalConnSnapshot(&message) + currentTransport := message.TransportConn + if currentTransport == nil && current != nil { + currentTransport = current.CurrentTransportConn() + } + if currentTransport != nil && !currentTransport.IsCurrent() { + if current != nil { + _ = s.replyDedicatedBulkAttach(current, message, toBulkAttachResponseError(newBulkAttachError(bulkAttachErrorCodeAttachFailed, true, transportDetachedErrorForTransport(currentTransport).Error()), "")) + } + return true + } var ( req bulkAttachRequest logical *LogicalConn @@ -854,6 +1064,9 @@ func (s *ServerCommon) resolveInboundDedicatedBulk(current *LogicalConn, req bul } } bulk.markDedicatedAttachAttempt() + if bulkTransport := bulk.TransportConn(); bulkTransport != nil && !bulkTransport.IsCurrent() { + return nil, nil, newBulkAttachError(bulkAttachErrorCodeAttachFailed, true, transportDetachedErrorForTransport(bulkTransport).Error()) + } if !bulk.Dedicated() { bulk.setDedicatedAttachLastCode(string(bulkAttachErrorCodeBulkNotDedicated)) return nil, nil, &bulkAttachError{ @@ -907,8 +1120,24 @@ func (s *ServerCommon) finishInboundDedicatedBulkAttach(current *LogicalConn, lo if current == nil || logical == nil || bulk == nil { return newBulkAttachError(bulkAttachErrorCodeInvalidRequest, false, errBulkLogicalConnNil.Error()) } + // Keep the target logical transport generation stable while the attach + // reply, sidecar publication, runtime rebinding and accept dispatch are + // committed. Reattach cleanup takes the same lock before retiring the old + // generation. + logical.transportLifecycleMu.Lock() + defer logical.transportLifecycleMu.Unlock() scope := serverFileScope(logical) laneID := bulk.dedicatedLaneIDSnapshot() + currentTransport := message.TransportConn + if currentTransport == nil { + currentTransport = current.CurrentTransportConn() + } + if currentTransport != nil && !currentTransport.IsCurrent() { + return newBulkAttachError(bulkAttachErrorCodeAttachFailed, true, transportDetachedErrorForTransport(currentTransport).Error()) + } + if bulkTransport := bulk.TransportConn(); bulkTransport != nil && !bulkTransport.IsCurrent() { + return newBulkAttachError(bulkAttachErrorCodeAttachFailed, true, transportDetachedErrorForTransport(bulkTransport).Error()) + } conn, err := current.detachTransportForTransfer() if err != nil { return newBulkAttachError(bulkAttachErrorCodeAttachFailed, true, err.Error()) @@ -962,6 +1191,11 @@ func (s *ServerCommon) finishInboundDedicatedBulkAttach(current *LogicalConn, lo bulk.setDedicatedAttachLastCode(string(bulkAttachErrorCodeAttachFailed)) return fail("bulk dedicated attach failed", err) } + if bulkTransport := bulk.TransportConn(); bulkTransport != nil && !bulkTransport.IsCurrent() { + sidecar.close() + bulk.setDedicatedAttachLastCode(string(bulkAttachErrorCodeAttachFailed)) + return fail("bulk dedicated attach transport replaced", transportDetachedErrorForTransport(bulkTransport)) + } if err := s.replyDedicatedBulkAttachDetached(current, conn, message, bulkAttachResponse{Accepted: true}); err != nil { bulk.setDedicatedAttachLastCode(string(bulkAttachErrorCodeAttachFailed)) sidecar.close() @@ -971,6 +1205,11 @@ func (s *ServerCommon) finishInboundDedicatedBulkAttach(current *LogicalConn, lo stopCurrent("bulk dedicated attach reply failed", err) return nil } + if bulkTransport := bulk.TransportConn(); bulkTransport != nil && !bulkTransport.IsCurrent() { + sidecar.close() + bulk.setDedicatedAttachLastCode(string(bulkAttachErrorCodeAttachFailed)) + return fail("bulk dedicated attach transport replaced", transportDetachedErrorForTransport(bulkTransport)) + } oldSidecar := s.installServerDedicatedSidecar(logical, laneID, sidecar) if runtime := s.getBulkRuntime(); runtime != nil { runtime.attachSharedDedicatedConn(scope, laneID, conn) @@ -985,7 +1224,7 @@ func (s *ServerCommon) finishInboundDedicatedBulkAttach(current *LogicalConn, lo oldSidecar.close() } go s.readDedicatedSidecarLoop(logical, sidecar) - s.startServerBulkAcceptDispatch(bulk, logical, messageTransportConnSnapshot(&message)) + s.startServerBulkAcceptDispatch(bulk, logical, bulk.TransportConn()) if runtime := s.getBulkRuntime(); runtime != nil { s.dispatchPendingServerBulkAccepts(scope, conn, bulk, logical) } @@ -1068,20 +1307,53 @@ func (s *ServerCommon) readDedicatedSidecarLoop(logical *LogicalConn, sidecar *b if s == nil || logical == nil || sidecar == nil || sidecar.conn == nil { return } - runtime := s.getBulkRuntime() + debug := s.IsDebugMode() scope := serverFileScope(logical) + sidecarCurrent := func() bool { + return s.serverDedicatedSidecarCurrent(logical, sidecar) + } for { + if !sidecarCurrent() { + if debug { + fmt.Printf("[bulk-debug] at=%s side=server lane=%d sidecar stopped reason=stale-connection\n", time.Now().Format(time.RFC3339Nano), sidecar.laneID) + } + return + } + var readStarted time.Time + if debug { + readStarted = time.Now() + } payload, payloadRelease, err := readBulkDedicatedRecordPooled(sidecar.conn) + if debug { + fmt.Printf("[bulk-debug] at=%s side=server lane=%d sidecar record read bytes=%d elapsed=%s error=%v\n", time.Now().Format(time.RFC3339Nano), sidecar.laneID, len(payload), time.Since(readStarted), err) + } if err != nil { + if !sidecarCurrent() { + return + } s.handleServerDedicatedSidecarFailure(logical, sidecar, err) return } + if !sidecarCurrent() { + if payloadRelease != nil { + payloadRelease() + } + return + } plain, plainRelease, err := decryptTransportPayloadCodecPooled(logical.protectionModeSnapshot(), logical.modernPSKRuntimeSnapshot(), logical.msgDeSnapshot(), logical.secretKeySnapshot(), payload, payloadRelease) if err != nil { + if !sidecarCurrent() { + return + } s.handleServerDedicatedSidecarFailure(logical, sidecar, err) return } owner := newBulkReadPayloadOwner(plainRelease) + if !sidecarCurrent() { + owner.done() + return + } + runtime := s.getBulkRuntime() if runtime == nil { owner.done() continue @@ -1090,15 +1362,39 @@ func (s *ServerCommon) readDedicatedSidecarLoop(logical *LogicalConn, sidecar *b currentDataID uint64 currentBulk *bulkHandle skipDataID bool + staleSidecar bool ) err = walkDedicatedBulkInboundPayload(plain, func(dataID uint64, item bulkDedicatedBatchItem) error { + if debug && item.Type == bulkFastPayloadTypeRelease { + fmt.Printf("[bulk-debug] at=%s side=server lane=%d data=%d sidecar release received payload=%d\n", time.Now().Format(time.RFC3339Nano), sidecar.laneID, dataID, len(item.Payload)) + } + if !sidecarCurrent() { + staleSidecar = true + currentBulk = nil + skipDataID = true + return nil + } if dataID != currentDataID { currentDataID = dataID currentBulk = nil skipDataID = false - bulk, ok := runtime.lookupByDataID(scope, dataID) + bulk, ok := runtime.lookupInboundFrame(scope, dataID) if !ok { - s.bestEffortRejectInboundDedicatedData(logical, sidecar.conn, dataID, errBulkNotFound.Error()) + if debug { + fmt.Printf("[bulk-debug] at=%s side=server id=- data=%d age=- sidecar inbound lookup failed type=%d\n", time.Now().Format(time.RFC3339Nano), dataID, item.Type) + } + if sidecarCurrent() { + s.bestEffortRejectInboundDedicatedData(logical, sidecar.conn, dataID, errBulkNotFound.Error()) + } else { + staleSidecar = true + } + skipDataID = true + return nil + } + if !bulkDedicatedSidecarConnCurrent(bulk, sidecar) { + if debug { + bulk.debugf("sidecar reject lane=%d type=%d reason=connection-mismatch", sidecar.laneID, item.Type) + } skipDataID = true return nil } @@ -1108,12 +1404,42 @@ func (s *ServerCommon) readDedicatedSidecarLoop(logical *LogicalConn, sidecar *b if skipDataID || currentBulk == nil { return nil } + if !sidecarCurrent() || !bulkDedicatedSidecarConnCurrent(currentBulk, sidecar) { + if !sidecarCurrent() { + staleSidecar = true + } + currentBulk = nil + skipDataID = true + return nil + } + var release func() if item.Type == bulkFastPayloadTypeData { release = owner.retainChunk() } + if !sidecarCurrent() || !bulkDedicatedSidecarConnCurrent(currentBulk, sidecar) { + if release != nil { + release() + } + if !sidecarCurrent() { + staleSidecar = true + } + currentBulk = nil + skipDataID = true + return nil + } dispatchErr := dispatchDedicatedBulkInboundItemWithRelease(currentBulk, item, release) + if dispatchErr != nil { + if debug { + currentBulk.debugf("sidecar dispatch failed lane=%d type=%d error=%v", sidecar.laneID, item.Type, dispatchErr) + } + if !sidecarCurrent() { + staleSidecar = true + currentBulk = nil + skipDataID = true + return nil + } if !errors.Is(dispatchErr, io.EOF) { _ = s.sendDedicatedBulkReset(context.Background(), logical, currentBulk, dispatchErr.Error()) currentBulk.markReset(dispatchErr) @@ -1129,13 +1455,17 @@ func (s *ServerCommon) readDedicatedSidecarLoop(logical *LogicalConn, sidecar *b return nil }) if err != nil { - if plainRelease != nil { - plainRelease() + owner.done() + if !sidecarCurrent() { + return } s.handleServerDedicatedSidecarFailure(logical, sidecar, err) return } owner.done() + if staleSidecar { + return + } } } @@ -1193,6 +1523,9 @@ func (c *ClientCommon) dedicatedBulkSender(bulk *bulkHandle) (*bulkDedicatedSend if actual != sender { sender.stop() } + if actual == nil { + return nil, io.ErrClosedPipe + } return actual, nil } @@ -1200,7 +1533,8 @@ func (c *ClientCommon) dedicatedBulkLaneSender(bulk *bulkHandle) (*bulkDedicated if c == nil || bulk == nil { return nil, errBulkClientNil } - sidecar := c.clientDedicatedSidecarSnapshotForLane(bulk.dedicatedLaneIDSnapshot()) + route := bulk.clientSessionRouteSnapshot() + sidecar := c.clientDedicatedSidecarSnapshotForLaneAtRoute(bulk.dedicatedLaneIDSnapshot(), route) conn := bulk.dedicatedConnSnapshot() if sidecar == nil || sidecar.conn == nil || conn == nil || sidecar.conn != conn { return nil, transportDetachedError("dedicated bulk sidecar not attached", nil) @@ -1215,7 +1549,7 @@ func (c *ClientCommon) dedicatedBulkLaneSender(bulk *bulkHandle) (*bulkDedicated return c.encodeDedicatedBulkBatchesPayloadPooledWithRuntime(laneRuntime, batches) }, func(err error) { c.handleClientDedicatedSidecarFailure(sidecar, err) - }) + }, c.IsDebugMode()) }) if sender == nil { return nil, transportDetachedError("dedicated bulk sidecar not attached", nil) @@ -1335,6 +1669,9 @@ func (s *ServerCommon) dedicatedBulkSender(logical *LogicalConn, bulk *bulkHandl if actual != sender { sender.stop() } + if actual == nil { + return nil, io.ErrClosedPipe + } return actual, nil } @@ -1362,7 +1699,7 @@ func (s *ServerCommon) dedicatedBulkLaneSender(logical *LogicalConn, bulk *bulkH return s.encodeDedicatedBulkBatchesPayloadPooledWithRuntime(logical, laneRuntime, batches) }, func(err error) { s.handleServerDedicatedSidecarFailure(logical, sidecar, err) - }) + }, s.IsDebugMode()) }) if sender == nil { return nil, transportDetachedError("dedicated bulk sidecar not attached", nil) @@ -1457,7 +1794,11 @@ func (c *ClientCommon) encodeDedicatedBulkBatchPayload(dataID uint64, items []bu } profile := c.clientTransportProtectionSnapshot() if runtime := profile.runtime; runtime != nil { - return runtime.sealFilledPayload(bulkDedicatedBatchPlainLen(items), func(dst []byte) error { + plainLen, err := bulkDedicatedBatchPlainLenChecked(items) + if err != nil { + return nil, err + } + return runtime.sealFilledPayload(plainLen, func(dst []byte) error { return writeBulkDedicatedBatchPlain(dst, dataID, items) }) } @@ -1486,7 +1827,11 @@ func (c *ClientCommon) encodeDedicatedBulkBatchesPayloadPooledWithRuntime(runtim return nil, nil, errBulkFastPayloadInvalid } if runtime != nil { - return runtime.sealFilledPayloadPooled(bulkDedicatedBatchesPlainLen(batches), func(dst []byte) error { + plainLen, err := bulkDedicatedBatchesPlainLenChecked(batches) + if err != nil { + return nil, nil, err + } + return runtime.sealFilledPayloadPooled(plainLen, func(dst []byte) error { return writeBulkDedicatedBatchesPlain(dst, batches) }) } @@ -1522,7 +1867,11 @@ func (s *ServerCommon) encodeDedicatedBulkBatchPayload(logical *LogicalConn, dat return nil, errBulkLogicalConnNil } if runtime := logical.modernPSKRuntimeSnapshot(); runtime != nil { - return runtime.sealFilledPayload(bulkDedicatedBatchPlainLen(items), func(dst []byte) error { + plainLen, err := bulkDedicatedBatchPlainLenChecked(items) + if err != nil { + return nil, err + } + return runtime.sealFilledPayload(plainLen, func(dst []byte) error { return writeBulkDedicatedBatchPlain(dst, dataID, items) }) } @@ -1547,7 +1896,11 @@ func (s *ServerCommon) encodeDedicatedBulkBatchesPayloadPooledWithRuntime(logica return nil, nil, errBulkFastPayloadInvalid } if runtime != nil { - return runtime.sealFilledPayloadPooled(bulkDedicatedBatchesPlainLen(batches), func(dst []byte) error { + plainLen, err := bulkDedicatedBatchesPlainLenChecked(batches) + if err != nil { + return nil, nil, err + } + return runtime.sealFilledPayloadPooled(plainLen, func(dst []byte) error { return writeBulkDedicatedBatchesPlain(dst, batches) }) } diff --git a/bulk_dedicated_attach_test.go b/bulk_dedicated_attach_test.go index ab6b2b8..8a275ea 100644 --- a/bulk_dedicated_attach_test.go +++ b/bulk_dedicated_attach_test.go @@ -185,6 +185,145 @@ func TestSendDedicatedBulkAttachRequestUsesBootstrapProtectionEvenAfterSteadySwi } } +func TestAttachDedicatedBulkSidecarStopsWhenOriginalRouteReattaches(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, ^uint32(0)) + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + epoch := client.beginClientSessionEpoch() + client.setClientSessionRuntime(newClientSessionRuntime(firstLeft, stopCtx, stopFn, queue, epoch)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + + route := client.clientSessionRouteSnapshot() + bulk := newBulkHandle(stopCtx, client.getBulkRuntime(), clientFileScope(), BulkOpenRequest{ + BulkID: "dedicated-original-route", + DataID: 1, + Dedicated: true, + DedicatedLaneID: 1, + AttachToken: "attach-token", + }, epoch, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.setClientSnapshotOwner(client) + bulk.setClientSessionRoute(route) + defer bulk.finalize() + + client.bulkDedicatedAttachSem = make(chan struct{}, 1) + client.bulkDedicatedAttachSem <- struct{}{} + defer func() { + select { + case <-client.bulkDedicatedAttachSem: + default: + } + }() + result := make(chan error, 1) + go func() { + result <- client.attachDedicatedBulkSidecar(context.Background(), bulk) + }() + + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + client.bulkDedicatedSidecarMu.Lock() + lane := client.bulkDedicatedLanes[1] + blocked := lane != nil && lane.attachFlight != nil + client.bulkDedicatedSidecarMu.Unlock() + if blocked { + break + } + time.Sleep(time.Millisecond) + } + client.bulkDedicatedSidecarMu.Lock() + lane := client.bulkDedicatedLanes[1] + blocked := lane != nil && lane.attachFlight != nil + client.bulkDedicatedSidecarMu.Unlock() + if !blocked { + t.Fatal("dedicated attach did not reach the blocked attach slot") + } + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := client.attachClientSessionTransport(secondLeft); err != nil { + t.Fatalf("attach replacement client transport: %v", err) + } + select { + case err := <-result: + if !errors.Is(err, errTransportDetached) { + t.Fatalf("dedicated attach error = %v, want transport detached", err) + } + case <-time.After(time.Second): + t.Fatal("dedicated attach remained blocked after original route detached") + } + if sidecar := client.clientDedicatedSidecarSnapshotForLane(1); sidecar != nil { + t.Fatalf("stale route installed dedicated sidecar: %+v", sidecar) + } +} + +func TestClientDedicatedSidecarRejectsLateOldRoutePublicationAndRelease(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, ^uint32(0)) + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + epoch := client.beginClientSessionEpoch() + client.setClientSessionRuntime(newClientSessionRuntime(firstLeft, stopCtx, stopFn, queue, epoch)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + + oldRoute := client.clientSessionRouteSnapshot() + oldLane, err := client.reserveBulkDedicatedLaneAtRoute(oldRoute) + if err != nil { + t.Fatalf("reserve old-route lane: %v", err) + } + oldSidecarLeft, oldSidecarRight := net.Pipe() + defer oldSidecarRight.Close() + oldSidecar := newBulkDedicatedSidecar(oldSidecarLeft, oldLane) + if _, installed, err := client.installClientDedicatedSidecarAtRoute(oldLane, oldSidecar, oldRoute); err != nil || !installed { + t.Fatalf("install old-route sidecar = installed=%v err=%v, want installed", installed, err) + } + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := client.attachClientSessionTransport(secondLeft); err != nil { + t.Fatalf("attach replacement client transport: %v", err) + } + newRoute := client.clientSessionRouteSnapshot() + lateLeft, lateRight := net.Pipe() + defer lateRight.Close() + lateSidecar := newBulkDedicatedSidecar(lateLeft, oldLane) + if _, installed, err := client.installClientDedicatedSidecarAtRoute(oldLane, lateSidecar, oldRoute); err == nil || installed { + lateSidecar.close() + t.Fatalf("late old-route sidecar install = installed=%v err=%v, want rejection", installed, err) + } + lateSidecar.close() + + newLane, err := client.reserveBulkDedicatedLaneAtRoute(newRoute) + if err != nil { + t.Fatalf("reserve new-route lane: %v", err) + } + newSidecarLeft, newSidecarRight := net.Pipe() + defer newSidecarRight.Close() + newSidecar := newBulkDedicatedSidecar(newSidecarLeft, newLane) + if _, installed, err := client.installClientDedicatedSidecarAtRoute(newLane, newSidecar, newRoute); err != nil || !installed { + t.Fatalf("install new-route sidecar = installed=%v err=%v, want installed", installed, err) + } + + client.releaseBulkDedicatedLaneAtRoute(newLane, oldRoute) + client.bulkDedicatedSidecarMu.Lock() + lane := client.bulkDedicatedLanes[newLane] + active := 0 + if lane != nil { + active = lane.activeBulks + } + client.bulkDedicatedSidecarMu.Unlock() + if active != 1 { + t.Fatalf("old-route release changed new-route lane active count to %d, want 1", active) + } +} + func TestHandleBulkAttachSystemMessageAcceptedWritesDirectReplyBeforeDedicatedHandoff(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) @@ -288,6 +427,112 @@ func TestHandleBulkAttachSystemMessageAcceptedWritesDirectReplyBeforeDedicatedHa } } +func TestHandleBulkAttachRejectsBulkWhosePrimaryTransportReattached(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) + defer runtimeCancel() + server.setServerSessionRuntime(&serverSessionRuntime{ + stopCtx: runtimeCtx, + stopFn: runtimeCancel, + queue: stario.NewQueueCtx(runtimeCtx, 4, ^uint32(0)), + }) + server.markSessionStarted() + defer server.markSessionStopped("test done", nil) + + attachLeft, attachRight := net.Pipe() + defer attachRight.Close() + current := server.bootstrapAcceptedLogical("dedicated-stale-primary-current", nil, attachLeft) + if current == nil { + t.Fatal("bootstrapAcceptedLogical(current) should return logical") + } + primaryLeft, primaryRight := net.Pipe() + defer primaryRight.Close() + target := server.bootstrapAcceptedLogical("dedicated-stale-primary-target", nil, primaryLeft) + if target == nil { + t.Fatal("bootstrapAcceptedLogical(target) should return logical") + } + primaryTransport := target.CurrentTransportConn() + if primaryTransport == nil { + t.Fatal("target primary transport should exist") + } + + bulk := newBulkHandle(context.Background(), server.getBulkRuntime(), serverFileScope(target), BulkOpenRequest{ + BulkID: "dedicated-stale-primary", + DataID: 7, + Dedicated: true, + AttachToken: "attach-token", + }, 0, target, primaryTransport, primaryTransport.TransportGeneration(), nil, nil, nil, nil, nil) + if err := server.getBulkRuntime().register(serverFileScope(target), bulk); err != nil { + t.Fatalf("register dedicated bulk: %v", err) + } + defer bulk.finalize() + + replacementLeft, replacementRight := net.Pipe() + defer replacementRight.Close() + if err := target.attachClientConnSessionTransport(replacementLeft); err != nil { + t.Fatalf("reattach target primary transport: %v", err) + } + reqPayload, err := server.sequenceEn(bulkAttachRequest{ + PeerID: target.ID(), + BulkID: bulk.ID(), + AttachToken: "attach-token", + }) + if err != nil { + t.Fatalf("encode bulk attach request: %v", err) + } + message := Message{ + NetType: NET_SERVER, + LogicalConn: current, + TransportConn: current.CurrentTransportConn(), + TransferMsg: TransferMsg{ + ID: 91, + Key: systemBulkAttachKey, + Value: reqPayload, + Type: MSG_SYS_WAIT, + }, + inboundConn: attachLeft, + } + + replyDone := make(chan bulkAttachResponse, 1) + go func() { + _ = attachRight.SetReadDeadline(time.Now().Add(time.Second)) + payload, readErr := readDirectSignalFramePayload(attachRight) + if readErr != nil { + replyDone <- bulkAttachResponse{Error: readErr.Error()} + return + } + transfer, decodeErr := decodeDirectSignalPayload(server.sequenceDe, current.msgDeSnapshot(), current.secretKeySnapshot(), payload) + if decodeErr != nil { + replyDone <- bulkAttachResponse{Error: decodeErr.Error()} + return + } + resp, decodeErr := decodeBulkAttachResponse(server.sequenceDe, transfer.Value) + if decodeErr != nil { + resp.Error = decodeErr.Error() + } + replyDone <- resp + }() + + if !server.handleBulkAttachSystemMessage(message) { + t.Fatal("handleBulkAttachSystemMessage should consume attach message") + } + select { + case resp := <-replyDone: + if resp.Accepted || resp.Error == "" { + t.Fatalf("stale-primary attach response = %+v, want rejection", resp) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for stale-primary attach rejection") + } + if !current.transportAttachedSnapshot() { + t.Fatal("rejected attach must not detach the inbound attach connection") + } + if got := bulk.dedicatedConnSnapshot(); got != nil { + t.Fatalf("rejected stale-primary bulk attached sidecar %v", got) + } +} + func TestHandleBulkAttachSystemMessageDoesNotExposeSharedSidecarBeforeReplyCompletes(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) diff --git a/bulk_dedicated_batch.go b/bulk_dedicated_batch.go index 5738c88..67cc2b9 100644 --- a/bulk_dedicated_batch.go +++ b/bulk_dedicated_batch.go @@ -75,12 +75,15 @@ type bulkDedicatedSender struct { encodeBatch func([]bulkDedicatedSendRequest) ([]byte, error) fail func(error) - reqCh chan bulkDedicatedBatchRequest - stopCh chan struct{} - doneCh chan struct{} - stopOnce sync.Once - flushMu sync.Mutex - queued atomic.Int64 + reqCh chan bulkDedicatedBatchRequest + stopCh chan struct{} + doneCh chan struct{} + stopOnce sync.Once + admissionMu sync.Mutex + admitting sync.WaitGroup + admissionClosed bool + flushMu sync.Mutex + queued atomic.Int64 errMu sync.Mutex err error @@ -220,19 +223,17 @@ func (s *bulkDedicatedSender) submitBatch(ctx context.Context, items []bulkDedic req.Ack = make(chan error, 1) } s.queued.Add(1) - select { - case <-ctx.Done(): + if !s.enqueue(req) { s.queued.Add(-1) - return normalizeStreamDeadlineError(ctx.Err()) - case <-s.stopCh: - s.queued.Add(-1) - return s.stoppedErr() - case s.reqCh <- req: - if !wait { - return nil + if err := ctx.Err(); err != nil { + return normalizeStreamDeadlineError(err) } - return s.waitAck(req) + return s.stoppedErr() } + if !wait { + return nil + } + return s.waitAck(req) } func (s *bulkDedicatedSender) tryDirectSubmitBatch(ctx context.Context, items []bulkDedicatedSendRequest) (bool, error) { @@ -276,9 +277,15 @@ func (s *bulkDedicatedSender) tryDirectSubmitBatch(ctx context.Context, items [] default: } deadline, _ := ctx.Deadline() + select { + case <-s.stopCh: + return true, s.stoppedErr() + default: + } if err := s.flush(items, deadline); err != nil { err = normalizeDedicatedBulkSendError(err) - s.setErr(err) + s.markFailed(err) + s.waitAdmissions() s.failPending(err) if s.fail != nil { go s.fail(err) @@ -311,11 +318,36 @@ func (s *bulkDedicatedSender) stop() { if s == nil { return } - s.stopOnce.Do(func() { - s.setErr(errTransportDetached) - close(s.stopCh) - }) + s.markFailed(errTransportDetached) + s.waitAdmissions() <-s.doneCh + // Direct submissions flush on the caller goroutine rather than run(). + // Wait for that path before allowing the underlying connection to close or + // be handed to a replacement sender. + s.flushMu.Lock() + s.flushMu.Unlock() +} + +func (s *bulkDedicatedSender) enqueue(req bulkDedicatedBatchRequest) bool { + if s == nil { + return false + } + s.admissionMu.Lock() + if s.admissionClosed { + s.admissionMu.Unlock() + return false + } + s.admitting.Add(1) + s.admissionMu.Unlock() + defer s.admitting.Done() + select { + case <-req.Ctx.Done(): + return false + case <-s.stopCh: + return false + case s.reqCh <- req: + return true + } } func (s *bulkDedicatedSender) run() { @@ -337,13 +369,19 @@ func (s *bulkDedicatedSender) run() { s.flushMu.Lock() err := s.errSnapshot() if err == nil { - err = s.flush(req.Items, req.Deadline) + select { + case <-s.stopCh: + err = s.stoppedErr() + default: + err = s.flush(req.Items, req.Deadline) + } } s.flushMu.Unlock() if err != nil { err = normalizeDedicatedBulkSendError(err) - s.setErr(err) + s.markFailed(err) s.finishRequest(req, err) + s.waitAdmissions() s.failPending(err) if s.fail != nil { go s.fail(err) @@ -390,6 +428,7 @@ func (r bulkDedicatedBatchRequest) canceledErr() error { func (s *bulkDedicatedSender) nextRequest() (bulkDedicatedBatchRequest, bool) { select { case <-s.stopCh: + s.waitAdmissions() s.failPending(s.stoppedErr()) return bulkDedicatedBatchRequest{}, false case req := <-s.reqCh: @@ -455,6 +494,25 @@ func (s *bulkDedicatedSender) setErr(err error) { s.errMu.Unlock() } +func (s *bulkDedicatedSender) markFailed(err error) { + if s == nil { + return + } + s.setErr(err) + s.stopOnce.Do(func() { + s.admissionMu.Lock() + s.admissionClosed = true + close(s.stopCh) + s.admissionMu.Unlock() + }) +} + +func (s *bulkDedicatedSender) waitAdmissions() { + if s != nil { + s.admitting.Wait() + } +} + func (s *bulkDedicatedSender) errSnapshot() error { if s == nil { return errTransportDetached @@ -525,11 +583,77 @@ func bulkDedicatedBatchesPlainLen(batches []bulkDedicatedOutboundBatch) int { } } -func encodeBulkDedicatedBatchesPlain(batches []bulkDedicatedOutboundBatch) ([]byte, error) { - if len(batches) == 0 { - return nil, errBulkFastPayloadInvalid +func bulkDedicatedBatchesPlainLenChecked(batches []bulkDedicatedOutboundBatch) (int, error) { + switch len(batches) { + case 0: + return 0, errBulkFastPayloadInvalid + case 1: + return bulkDedicatedBatchPlainLenChecked(batches[0].Items) + } + if len(batches) > bulkDedicatedBatchMaxItems { + return 0, errBulkFastPayloadInvalid + } + total := bulkDedicatedSuperBatchHeaderLen + totalItems := 0 + for _, batch := range batches { + if batch.DataID == 0 || len(batch.Items) == 0 || len(batch.Items) > bulkDedicatedBatchMaxItems { + return 0, errBulkFastPayloadInvalid + } + totalItems += len(batch.Items) + if totalItems > bulkDedicatedBatchMaxItems { + return 0, errBulkFastPayloadInvalid + } + if total > bulkDedicatedBatchMaxPlainBytes-bulkDedicatedSuperBatchGroupHeaderLen { + return 0, errBulkFastPayloadInvalid + } + total += bulkDedicatedSuperBatchGroupHeaderLen + itemsLen, err := bulkDedicatedSendRequestsLenChecked(batch.Items) + if err != nil { + return 0, err + } + if total > bulkDedicatedBatchMaxPlainBytes-itemsLen { + return 0, errBulkFastPayloadInvalid + } + total += itemsLen + } + return total, nil +} + +func bulkDedicatedBatchPlainLenChecked(items []bulkDedicatedSendRequest) (int, error) { + if len(items) == 0 || len(items) > bulkDedicatedBatchMaxItems { + return 0, errBulkFastPayloadInvalid + } + total := bulkDedicatedBatchHeaderLen + itemsLen, err := bulkDedicatedSendRequestsLenChecked(items) + if err != nil { + return 0, err + } + if total > bulkDedicatedBatchMaxPlainBytes-itemsLen { + return 0, errBulkFastPayloadInvalid + } + return total + itemsLen, nil +} + +func bulkDedicatedSendRequestsLenChecked(items []bulkDedicatedSendRequest) (int, error) { + total := 0 + for _, item := range items { + itemLen := bulkDedicatedSendRequestLen(item) + if itemLen < bulkDedicatedBatchItemHeaderLen || itemLen > bulkDedicatedBatchMaxPlainBytes { + return 0, errBulkFastPayloadInvalid + } + if total > bulkDedicatedBatchMaxPlainBytes-itemLen { + return 0, errBulkFastPayloadInvalid + } + total += itemLen + } + return total, nil +} + +func encodeBulkDedicatedBatchesPlain(batches []bulkDedicatedOutboundBatch) ([]byte, error) { + total, err := bulkDedicatedBatchesPlainLenChecked(batches) + if err != nil { + return nil, err } - total := bulkDedicatedBatchesPlainLen(batches) buf := make([]byte, total) if err := writeBulkDedicatedBatchesPlain(buf, batches); err != nil { return nil, err @@ -552,7 +676,11 @@ func writeBulkDedicatedSuperBatchPlain(buf []byte, batches []bulkDedicatedOutbou if len(batches) <= 1 { return errBulkFastPayloadInvalid } - if len(buf) != bulkDedicatedBatchesPlainLen(batches) || len(buf) > bulkDedicatedBatchMaxPlainBytes { + plainLen, err := bulkDedicatedBatchesPlainLenChecked(batches) + if err != nil { + return err + } + if len(buf) != plainLen { return errBulkFastPayloadInvalid } copy(buf[:4], bulkDedicatedSuperBatchMagic) @@ -630,7 +758,10 @@ func encodeBulkDedicatedBatchesPayloadFast(encode transportFastPlainEncoder, sec if encode == nil { return nil, errTransportPayloadEncryptFailed } - plainLen := bulkDedicatedBatchesPlainLen(batches) + plainLen, err := bulkDedicatedBatchesPlainLenChecked(batches) + if err != nil { + return nil, err + } return encode(secretKey, plainLen, func(dst []byte) error { return writeBulkDedicatedBatchesPlain(dst, batches) }) @@ -645,10 +776,14 @@ func bulkDedicatedBatchPlainLen(items []bulkDedicatedSendRequest) int { } func writeBulkDedicatedBatchPlain(buf []byte, dataID uint64, items []bulkDedicatedSendRequest) error { - if dataID == 0 || len(items) == 0 { + if dataID == 0 { return errBulkFastPayloadInvalid } - if len(buf) != bulkDedicatedBatchPlainLen(items) { + plainLen, err := bulkDedicatedBatchPlainLenChecked(items) + if err != nil { + return err + } + if len(buf) != plainLen { return errBulkFastPayloadInvalid } copy(buf[:4], bulkDedicatedBatchMagic) @@ -679,10 +814,11 @@ func decodeBulkDedicatedBatchPlain(payload []byte) (uint64, []bulkDedicatedBatch return 0, nil, true, errBulkFastPayloadInvalid } dataID := binary.BigEndian.Uint64(payload[8:16]) - count := int(binary.BigEndian.Uint32(payload[16:20])) - if dataID == 0 || count <= 0 { + wireCount := binary.BigEndian.Uint32(payload[16:20]) + if dataID == 0 || wireCount == 0 || wireCount > bulkDedicatedBatchMaxItems { return 0, nil, true, errBulkFastPayloadInvalid } + count := int(wireCount) items := make([]bulkDedicatedBatchItem, 0, count) offset := bulkDedicatedBatchHeaderLen for i := 0; i < count; i++ { @@ -697,11 +833,12 @@ func decodeBulkDedicatedBatchPlain(payload []byte) (uint64, []bulkDedicatedBatch } flags := payload[offset+1] seq := binary.BigEndian.Uint64(payload[offset+4 : offset+12]) - dataLen := int(binary.BigEndian.Uint32(payload[offset+12 : offset+16])) + wireDataLen := binary.BigEndian.Uint32(payload[offset+12 : offset+16]) offset += bulkDedicatedBatchItemHeaderLen - if dataLen < 0 || len(payload)-offset < dataLen { + if uint64(wireDataLen) > uint64(len(payload)-offset) { return 0, nil, true, errBulkFastPayloadInvalid } + dataLen := int(wireDataLen) items = append(items, bulkDedicatedBatchItem{ Type: itemType, Flags: flags, @@ -726,10 +863,11 @@ func decodeBulkDedicatedSuperBatchPlain(payload []byte) ([]bulkDedicatedInboundB if payload[4] != bulkDedicatedSuperBatchVersion { return nil, true, errBulkFastPayloadInvalid } - groupCount := int(binary.BigEndian.Uint32(payload[8:12])) - if groupCount <= 0 { + wireGroupCount := binary.BigEndian.Uint32(payload[8:12]) + if wireGroupCount == 0 || wireGroupCount > bulkDedicatedBatchMaxItems { return nil, true, errBulkFastPayloadInvalid } + groupCount := int(wireGroupCount) batches := make([]bulkDedicatedInboundBatch, 0, groupCount) offset := bulkDedicatedSuperBatchHeaderLen totalItems := 0 @@ -738,11 +876,12 @@ func decodeBulkDedicatedSuperBatchPlain(payload []byte) ([]bulkDedicatedInboundB return nil, true, errBulkFastPayloadInvalid } dataID := binary.BigEndian.Uint64(payload[offset : offset+8]) - count := int(binary.BigEndian.Uint32(payload[offset+8 : offset+12])) + wireCount := binary.BigEndian.Uint32(payload[offset+8 : offset+12]) offset += bulkDedicatedSuperBatchGroupHeaderLen - if dataID == 0 || count <= 0 { + if dataID == 0 || wireCount == 0 || wireCount > bulkDedicatedBatchMaxItems { return nil, true, errBulkFastPayloadInvalid } + count := int(wireCount) totalItems += count if totalItems > bulkDedicatedBatchMaxItems { return nil, true, errBulkFastPayloadInvalid @@ -760,11 +899,12 @@ func decodeBulkDedicatedSuperBatchPlain(payload []byte) ([]bulkDedicatedInboundB } flags := payload[offset+1] seq := binary.BigEndian.Uint64(payload[offset+4 : offset+12]) - dataLen := int(binary.BigEndian.Uint32(payload[offset+12 : offset+16])) + wireDataLen := binary.BigEndian.Uint32(payload[offset+12 : offset+16]) offset += bulkDedicatedBatchItemHeaderLen - if dataLen < 0 || len(payload)-offset < dataLen { + if uint64(wireDataLen) > uint64(len(payload)-offset) { return nil, true, errBulkFastPayloadInvalid } + dataLen := int(wireDataLen) items = append(items, bulkDedicatedBatchItem{ Type: itemType, Flags: flags, @@ -881,10 +1021,11 @@ func walkDedicatedBulkInboundBatchPlain(payload []byte, visit func(dataID uint64 return errBulkFastPayloadInvalid } dataID := binary.BigEndian.Uint64(payload[8:16]) - count := int(binary.BigEndian.Uint32(payload[16:20])) - if dataID == 0 || count <= 0 { + wireCount := binary.BigEndian.Uint32(payload[16:20]) + if dataID == 0 || wireCount == 0 || wireCount > bulkDedicatedBatchMaxItems { return errBulkFastPayloadInvalid } + count := int(wireCount) offset := bulkDedicatedBatchHeaderLen for i := 0; i < count; i++ { if len(payload)-offset < bulkDedicatedBatchItemHeaderLen { @@ -898,11 +1039,12 @@ func walkDedicatedBulkInboundBatchPlain(payload []byte, visit func(dataID uint64 } flags := payload[offset+1] seq := binary.BigEndian.Uint64(payload[offset+4 : offset+12]) - dataLen := int(binary.BigEndian.Uint32(payload[offset+12 : offset+16])) + wireDataLen := binary.BigEndian.Uint32(payload[offset+12 : offset+16]) offset += bulkDedicatedBatchItemHeaderLen - if dataLen < 0 || len(payload)-offset < dataLen { + if uint64(wireDataLen) > uint64(len(payload)-offset) { return errBulkFastPayloadInvalid } + dataLen := int(wireDataLen) if err := visit(dataID, bulkDedicatedBatchItem{ Type: itemType, Flags: flags, @@ -926,10 +1068,11 @@ func walkDedicatedBulkInboundSuperBatchPlain(payload []byte, visit func(dataID u if payload[4] != bulkDedicatedSuperBatchVersion { return errBulkFastPayloadInvalid } - groupCount := int(binary.BigEndian.Uint32(payload[8:12])) - if groupCount <= 0 { + wireGroupCount := binary.BigEndian.Uint32(payload[8:12]) + if wireGroupCount == 0 || wireGroupCount > bulkDedicatedBatchMaxItems { return errBulkFastPayloadInvalid } + groupCount := int(wireGroupCount) offset := bulkDedicatedSuperBatchHeaderLen totalItems := 0 for i := 0; i < groupCount; i++ { @@ -937,11 +1080,12 @@ func walkDedicatedBulkInboundSuperBatchPlain(payload []byte, visit func(dataID u return errBulkFastPayloadInvalid } dataID := binary.BigEndian.Uint64(payload[offset : offset+8]) - count := int(binary.BigEndian.Uint32(payload[offset+8 : offset+12])) + wireCount := binary.BigEndian.Uint32(payload[offset+8 : offset+12]) offset += bulkDedicatedSuperBatchGroupHeaderLen - if dataID == 0 || count <= 0 { + if dataID == 0 || wireCount == 0 || wireCount > bulkDedicatedBatchMaxItems { return errBulkFastPayloadInvalid } + count := int(wireCount) totalItems += count if totalItems > bulkDedicatedBatchMaxItems { return errBulkFastPayloadInvalid @@ -958,11 +1102,12 @@ func walkDedicatedBulkInboundSuperBatchPlain(payload []byte, visit func(dataID u } flags := payload[offset+1] seq := binary.BigEndian.Uint64(payload[offset+4 : offset+12]) - dataLen := int(binary.BigEndian.Uint32(payload[offset+12 : offset+16])) + wireDataLen := binary.BigEndian.Uint32(payload[offset+12 : offset+16]) offset += bulkDedicatedBatchItemHeaderLen - if dataLen < 0 || len(payload)-offset < dataLen { + if uint64(wireDataLen) > uint64(len(payload)-offset) { return errBulkFastPayloadInvalid } + dataLen := int(wireDataLen) if err := visit(dataID, bulkDedicatedBatchItem{ Type: itemType, Flags: flags, diff --git a/bulk_dedicated_lane_sender.go b/bulk_dedicated_lane_sender.go index 2608dc4..1e5c06c 100644 --- a/bulk_dedicated_lane_sender.go +++ b/bulk_dedicated_lane_sender.go @@ -2,6 +2,7 @@ package notify import ( "context" + "fmt" "net" "sync" "sync/atomic" @@ -28,22 +29,27 @@ type bulkDedicatedLaneSender struct { encode func([]bulkDedicatedOutboundBatch) ([]byte, func(), error) fail func(error) - reqCh chan *bulkDedicatedLaneBatchRequest - stopCh chan struct{} - doneCh chan struct{} - stopOnce sync.Once - flushMu sync.Mutex - queued atomic.Int64 + reqCh chan *bulkDedicatedLaneBatchRequest + stopCh chan struct{} + doneCh chan struct{} + stopOnce sync.Once + admissionMu sync.Mutex + admitting sync.WaitGroup + admissionClosed bool + flushMu sync.Mutex + queued atomic.Int64 + debug bool errMu sync.Mutex err error } -func newBulkDedicatedLaneSender(conn net.Conn, encode func([]bulkDedicatedOutboundBatch) ([]byte, func(), error), fail func(error)) *bulkDedicatedLaneSender { +func newBulkDedicatedLaneSender(conn net.Conn, encode func([]bulkDedicatedOutboundBatch) ([]byte, func(), error), fail func(error), debug ...bool) *bulkDedicatedLaneSender { sender := &bulkDedicatedLaneSender{ conn: conn, encode: encode, fail: fail, + debug: len(debug) > 0 && debug[0], reqCh: make(chan *bulkDedicatedLaneBatchRequest, bulkDedicatedSendQueueSize), stopCh: make(chan struct{}), doneCh: make(chan struct{}), @@ -210,7 +216,16 @@ func (s *bulkDedicatedLaneSender) submitControl(ctx context.Context, dataID uint if len(payload) > 0 { items[0].Payload = append([]byte(nil), payload...) } - return s.submitBatch(ctx, dataID, items, true, false) + var startedAt time.Time + if s.debug && frameType == bulkFastPayloadTypeRelease { + startedAt = time.Now() + fmt.Printf("[bulk-debug] at=%s lane release enqueue data=%d queued=%d\n", startedAt.Format(time.RFC3339Nano), dataID, s.queued.Load()) + } + err := s.submitBatch(ctx, dataID, items, true, false) + if s.debug && frameType == bulkFastPayloadTypeRelease { + fmt.Printf("[bulk-debug] at=%s lane release done data=%d elapsed=%s error=%v\n", time.Now().Format(time.RFC3339Nano), dataID, time.Since(startedAt), err) + } + return err } func (s *bulkDedicatedLaneSender) submitBatch(ctx context.Context, dataID uint64, items []bulkDedicatedSendRequest, wait bool, borrowItems bool) error { @@ -226,21 +241,18 @@ func (s *bulkDedicatedLaneSender) submitBatch(ctx context.Context, dataID uint64 req := getBulkDedicatedLaneBatchRequest() req.prepare(ctx, dataID, items, wait, borrowItems) s.queued.Add(1) - select { - case <-ctx.Done(): + if !s.enqueue(req) { s.queued.Add(-1) req.recycle() - return normalizeStreamDeadlineError(ctx.Err()) - case <-s.stopCh: - s.queued.Add(-1) - req.recycle() - return s.stoppedErr() - case s.reqCh <- req: - if !wait { - return nil + if err := ctx.Err(); err != nil { + return normalizeStreamDeadlineError(err) } - return s.waitAck(req) + return s.stoppedErr() } + if !wait { + return nil + } + return s.waitAck(req) } func (s *bulkDedicatedLaneSender) tryDirectSubmitWrite(ctx context.Context, dataID uint64, startSeq uint64, payload []byte, chunkSize int) (bool, int, error) { @@ -325,12 +337,18 @@ func (s *bulkDedicatedLaneSender) tryDirectSubmitWrite(ctx context.Context, data seq++ written = end } + select { + case <-s.stopCh: + return true, start, s.stoppedErr() + default: + } if err := s.flush([]bulkDedicatedOutboundBatch{{ DataID: dataID, Items: items, }}, deadline); err != nil { err = normalizeDedicatedBulkSendError(err) - s.setErr(err) + s.markFailed(err) + s.waitAdmissions() s.failPending(err) if s.fail != nil { go s.fail(err) @@ -382,12 +400,18 @@ func (s *bulkDedicatedLaneSender) tryDirectSubmitBatch(ctx context.Context, data default: } deadline, _ := ctx.Deadline() + select { + case <-s.stopCh: + return true, s.stoppedErr() + default: + } if err := s.flush([]bulkDedicatedOutboundBatch{{ DataID: dataID, Items: items, }}, deadline); err != nil { err = normalizeDedicatedBulkSendError(err) - s.setErr(err) + s.markFailed(err) + s.waitAdmissions() s.failPending(err) if s.fail != nil { go s.fail(err) @@ -427,11 +451,35 @@ func (s *bulkDedicatedLaneSender) stop() { if s == nil { return } - s.stopOnce.Do(func() { - s.setErr(errTransportDetached) - close(s.stopCh) - }) + s.markFailed(errTransportDetached) + s.waitAdmissions() <-s.doneCh + // Direct submissions flush on the caller goroutine rather than run(). + // Drain that path before the connection can be handed to a replacement. + s.flushMu.Lock() + s.flushMu.Unlock() +} + +func (s *bulkDedicatedLaneSender) enqueue(req *bulkDedicatedLaneBatchRequest) bool { + if s == nil || req == nil { + return false + } + s.admissionMu.Lock() + if s.admissionClosed { + s.admissionMu.Unlock() + return false + } + s.admitting.Add(1) + s.admissionMu.Unlock() + defer s.admitting.Done() + select { + case <-req.Ctx.Done(): + return false + case <-s.stopCh: + return false + case s.reqCh <- req: + return true + } } func (s *bulkDedicatedLaneSender) run() { @@ -456,25 +504,42 @@ func (s *bulkDedicatedLaneSender) run() { DataID: req.DataID, Items: req.Items, }} - batchBytes := bulkDedicatedBatchesPlainLen(batches) + batchBytes, err := bulkDedicatedBatchesPlainLenChecked(batches) + if err != nil { + s.finishRequest(req, err) + continue + } deadline := req.Deadline + var lockStarted time.Time + if s.debug { + lockStarted = time.Now() + } s.flushMu.Lock() - err := s.errSnapshot() + if s.debug { + fmt.Printf("[bulk-debug] at=%s lane flush acquired data=%d lock-wait=%s queued=%d\n", time.Now().Format(time.RFC3339Nano), req.DataID, time.Since(lockStarted), s.queued.Load()) + } + err = s.errSnapshot() if err == nil { carry, err = s.collectBatchRequests(&batchReqs, &batches, &batchBytes, &deadline) if err == nil { - err = s.flush(batches, deadline) + select { + case <-s.stopCh: + err = s.stoppedErr() + default: + err = s.flush(batches, deadline) + } } } s.flushMu.Unlock() if err != nil { err = normalizeDedicatedBulkSendError(err) - s.setErr(err) + s.markFailed(err) s.finishBatchRequests(batchReqs, err) if carry != nil { s.finishRequest(carry, err) carry = nil } + s.waitAdmissions() s.failPending(err) if s.fail != nil { go s.fail(err) @@ -566,6 +631,7 @@ func (s *bulkDedicatedLaneSender) nextRequest(carry *bulkDedicatedLaneBatchReque select { case <-s.stopCh: err := s.stoppedErr() + s.waitAdmissions() s.finishRequest(carry, err) s.failPending(err) return nil, false @@ -575,6 +641,7 @@ func (s *bulkDedicatedLaneSender) nextRequest(carry *bulkDedicatedLaneBatchReque } select { case <-s.stopCh: + s.waitAdmissions() s.failPending(s.stoppedErr()) return nil, false case req := <-s.reqCh: @@ -670,6 +737,10 @@ func (s *bulkDedicatedLaneSender) flush(batches []bulkDedicatedOutboundBatch, de if s == nil || s.conn == nil { return errTransportDetached } + var startedAt time.Time + if s.debug { + startedAt = time.Now() + } payload, release, err := s.encode(batches) if err != nil { return err @@ -677,7 +748,19 @@ func (s *bulkDedicatedLaneSender) flush(batches []bulkDedicatedOutboundBatch, de if release != nil { defer release() } - return writeBulkDedicatedRecordWithDeadline(s.conn, payload, deadline) + if !s.debug { + return writeBulkDedicatedRecordWithDeadline(s.conn, payload, deadline) + } + encodeElapsed := time.Since(startedAt) + dataID := uint64(0) + if len(batches) > 0 { + dataID = batches[0].DataID + } + fmt.Printf("[bulk-debug] at=%s lane write begin data=%d groups=%d bytes=%d encode=%s\n", time.Now().Format(time.RFC3339Nano), dataID, len(batches), len(payload), encodeElapsed) + writeStarted := time.Now() + err = writeBulkDedicatedRecordWithDeadlineTrace(s.conn, payload, deadline, dataID) + fmt.Printf("[bulk-debug] at=%s lane write end data=%d elapsed=%s error=%v\n", time.Now().Format(time.RFC3339Nano), dataID, time.Since(writeStarted), err) + return err } func (s *bulkDedicatedLaneSender) finishRequest(req *bulkDedicatedLaneBatchRequest, err error) { @@ -726,6 +809,25 @@ func (s *bulkDedicatedLaneSender) setErr(err error) { s.errMu.Unlock() } +func (s *bulkDedicatedLaneSender) markFailed(err error) { + if s == nil { + return + } + s.setErr(err) + s.stopOnce.Do(func() { + s.admissionMu.Lock() + s.admissionClosed = true + close(s.stopCh) + s.admissionMu.Unlock() + }) +} + +func (s *bulkDedicatedLaneSender) waitAdmissions() { + if s != nil { + s.admitting.Wait() + } +} + func (s *bulkDedicatedLaneSender) errSnapshot() error { if s == nil { return errTransportDetached diff --git a/bulk_dedicated_sidecar.go b/bulk_dedicated_sidecar.go index 3d4627b..e389dc9 100644 --- a/bulk_dedicated_sidecar.go +++ b/bulk_dedicated_sidecar.go @@ -10,6 +10,8 @@ type bulkDedicatedSidecar struct { laneID uint32 conn net.Conn closeOnce sync.Once + connMu sync.Mutex + closed bool senderMu sync.Mutex sender *bulkDedicatedLaneSender @@ -17,6 +19,7 @@ type bulkDedicatedSidecar struct { type bulkDedicatedLane struct { id uint32 + route clientSessionRoute activeBulks int sidecar *bulkDedicatedSidecar attachFlight *bulkDedicatedAttachFlight @@ -81,8 +84,12 @@ func (s *bulkDedicatedSidecar) close() { return } s.closeOnce.Do(func() { - if s.conn != nil { - _ = s.conn.Close() + s.connMu.Lock() + s.closed = true + conn := s.conn + s.connMu.Unlock() + if conn != nil { + _ = conn.Close() } if sender := s.laneSenderSnapshot(); sender != nil { sender.stop() @@ -90,6 +97,21 @@ func (s *bulkDedicatedSidecar) close() { }) } +// withConn serializes logical attachment with sidecar cleanup. Without this +// boundary a cleanup can close the socket after lookup but before the bulk +// handle installs it. +func (s *bulkDedicatedSidecar) withConn(fn func(net.Conn) error) error { + if s == nil || fn == nil { + return errTransportDetached + } + s.connMu.Lock() + defer s.connMu.Unlock() + if s.closed || s.conn == nil { + return errTransportDetached + } + return fn(s.conn) +} + func (s *bulkDedicatedSidecar) laneSenderSnapshot() *bulkDedicatedLaneSender { if s == nil { return nil @@ -108,10 +130,13 @@ func (s *bulkDedicatedSidecar) laneSenderWithFactory(factory func(net.Conn) *bul if s.sender != nil { return s.sender } - if s.conn == nil { + s.connMu.Lock() + if s.closed || s.conn == nil { + s.connMu.Unlock() return nil } s.sender = factory(s.conn) + s.connMu.Unlock() return s.sender } @@ -119,18 +144,13 @@ func (c *ClientCommon) clientDedicatedSidecarSnapshot() *bulkDedicatedSidecar { if c == nil { return nil } + route := c.clientSessionRouteSnapshot() c.bulkDedicatedSidecarMu.Lock() defer c.bulkDedicatedSidecarMu.Unlock() - return firstClientDedicatedSidecarLocked(c.bulkDedicatedLanes) -} - -func firstClientDedicatedSidecarLocked(lanes map[uint32]*bulkDedicatedLane) *bulkDedicatedSidecar { - var ( - selected *bulkDedicatedSidecar - bestID uint32 - ) - for laneID, lane := range lanes { - if lane == nil || lane.sidecar == nil { + var selected *bulkDedicatedSidecar + var bestID uint32 + for laneID, lane := range c.bulkDedicatedLanes { + if lane == nil || lane.sidecar == nil || !sameClientDedicatedLaneRoute(lane, route) { continue } if selected == nil || laneID < bestID { @@ -142,18 +162,31 @@ func firstClientDedicatedSidecarLocked(lanes map[uint32]*bulkDedicatedLane) *bul } func (c *ClientCommon) reserveBulkDedicatedLane() uint32 { + laneID, _ := c.reserveBulkDedicatedLaneAtRoute(c.clientSessionRouteSnapshot()) + return laneID +} + +func (c *ClientCommon) reserveBulkDedicatedLaneAtRoute(route clientSessionRoute) (uint32, error) { if c == nil { - return normalizeBulkDedicatedLaneID(0) + return 0, errBulkClientNil + } + if route.bound() { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return 0, err + } } c.bulkDedicatedSidecarMu.Lock() - defer c.bulkDedicatedSidecarMu.Unlock() + if route.bound() && !c.clientSessionRouteCurrent(route) { + c.bulkDedicatedSidecarMu.Unlock() + return 0, transportDetachedSessionEpochError() + } if c.bulkDedicatedLanes == nil { c.bulkDedicatedLanes = make(map[uint32]*bulkDedicatedLane) } limit := c.bulkDedicatedLaneLimitSnapshot() var best *bulkDedicatedLane for _, lane := range c.bulkDedicatedLanes { - if lane == nil { + if lane == nil || !sameClientDedicatedLaneRoute(lane, route) { continue } if best == nil || lane.activeBulks < best.activeBulks || (lane.activeBulks == best.activeBulks && lane.id < best.id) { @@ -161,16 +194,73 @@ func (c *ClientCommon) reserveBulkDedicatedLane() uint32 { } } if best == nil || ((limit <= 0 || len(c.bulkDedicatedLanes) < limit) && best.activeBulks > 0) { - c.bulkDedicatedNextLaneID++ - laneID := normalizeBulkDedicatedLaneID(c.bulkDedicatedNextLaneID) - best = &bulkDedicatedLane{id: laneID} + laneID := c.bulkDedicatedNextLaneID + for { + laneID++ + laneID = normalizeBulkDedicatedLaneID(laneID) + if _, exists := c.bulkDedicatedLanes[laneID]; !exists { + break + } + } + c.bulkDedicatedNextLaneID = laneID + best = &bulkDedicatedLane{id: laneID, route: route} c.bulkDedicatedLanes[laneID] = best } best.activeBulks++ - return best.id + c.bulkDedicatedSidecarMu.Unlock() + return best.id, nil +} + +func (c *ClientCommon) retainBulkDedicatedLane(laneID uint32) uint32 { + _ = c.retainBulkDedicatedLaneAtRoute(laneID, c.clientSessionRouteSnapshot()) + return normalizeBulkDedicatedLaneID(laneID) +} + +func (c *ClientCommon) retainBulkDedicatedLaneAtRoute(laneID uint32, route clientSessionRoute) error { + laneID = normalizeBulkDedicatedLaneID(laneID) + if c == nil { + return errBulkClientNil + } + if route.bound() { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return err + } + } + c.bulkDedicatedSidecarMu.Lock() + if route.bound() && !c.clientSessionRouteCurrent(route) { + c.bulkDedicatedSidecarMu.Unlock() + return transportDetachedSessionEpochError() + } + if c.bulkDedicatedLanes == nil { + c.bulkDedicatedLanes = make(map[uint32]*bulkDedicatedLane) + } + lane := c.bulkDedicatedLanes[laneID] + var retiredSidecar *bulkDedicatedSidecar + var retiredFlight *bulkDedicatedAttachFlight + if lane == nil || !sameClientDedicatedLaneRoute(lane, route) { + if lane != nil { + retiredSidecar = lane.sidecar + retiredFlight = lane.attachFlight + } + lane = &bulkDedicatedLane{id: laneID, route: route} + c.bulkDedicatedLanes[laneID] = lane + } + lane.activeBulks++ + c.bulkDedicatedSidecarMu.Unlock() + if retiredSidecar != nil { + retiredSidecar.close() + } + if retiredFlight != nil { + retiredFlight.finish(errTransportDetached) + } + return nil } func (c *ClientCommon) releaseBulkDedicatedLane(laneID uint32) { + c.releaseBulkDedicatedLaneAtRoute(laneID, c.clientSessionRouteSnapshot()) +} + +func (c *ClientCommon) releaseBulkDedicatedLaneAtRoute(laneID uint32, route clientSessionRoute) { if c == nil { return } @@ -178,7 +268,7 @@ func (c *ClientCommon) releaseBulkDedicatedLane(laneID uint32) { c.bulkDedicatedSidecarMu.Lock() defer c.bulkDedicatedSidecarMu.Unlock() lane := c.bulkDedicatedLanes[laneID] - if lane == nil { + if lane == nil || !sameClientDedicatedLaneRoute(lane, route) { return } if lane.activeBulks > 0 { @@ -189,43 +279,105 @@ func (c *ClientCommon) releaseBulkDedicatedLane(laneID uint32) { } } +func sameClientDedicatedLaneRoute(lane *bulkDedicatedLane, route clientSessionRoute) bool { + if lane == nil { + return false + } + if lane.route.bound() || route.bound() { + return sameClientSessionRoute(lane.route, route) + } + return true +} + func (c *ClientCommon) clientDedicatedSidecarSnapshotForLane(laneID uint32) *bulkDedicatedSidecar { if c == nil { return nil } + route := c.clientSessionRouteSnapshot() + return c.clientDedicatedSidecarSnapshotForLaneAtRoute(laneID, route) +} + +func (c *ClientCommon) clientDedicatedSidecarSnapshotForLaneAtRoute(laneID uint32, route clientSessionRoute) *bulkDedicatedSidecar { + if c == nil { + return nil + } + if route.bound() && !c.clientSessionRouteCurrent(route) { + return nil + } laneID = normalizeBulkDedicatedLaneID(laneID) c.bulkDedicatedSidecarMu.Lock() defer c.bulkDedicatedSidecarMu.Unlock() + if route.bound() && !c.clientSessionRouteCurrent(route) { + return nil + } if lane := c.bulkDedicatedLanes[laneID]; lane != nil { + if !sameClientDedicatedLaneRoute(lane, route) { + return nil + } return lane.sidecar } return nil } -func (c *ClientCommon) beginClientDedicatedSidecarAttach(laneID uint32) (*bulkDedicatedSidecar, *bulkDedicatedAttachFlight, bool) { +func (c *ClientCommon) beginClientDedicatedSidecarAttach(laneID uint32, route clientSessionRoute) (*bulkDedicatedSidecar, *bulkDedicatedAttachFlight, bool, error) { if c == nil { - return nil, nil, false + return nil, nil, false, errBulkClientNil + } + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return nil, nil, false, err } laneID = normalizeBulkDedicatedLaneID(laneID) c.bulkDedicatedSidecarMu.Lock() - defer c.bulkDedicatedSidecarMu.Unlock() + if !c.clientSessionRouteCurrent(route) { + c.bulkDedicatedSidecarMu.Unlock() + return nil, nil, false, transportDetachedSessionEpochError() + } if c.bulkDedicatedLanes == nil { c.bulkDedicatedLanes = make(map[uint32]*bulkDedicatedLane) } lane := c.bulkDedicatedLanes[laneID] - if lane == nil { - lane = &bulkDedicatedLane{id: laneID} + var retiredSidecar *bulkDedicatedSidecar + var retiredFlight *bulkDedicatedAttachFlight + if lane == nil || !sameClientDedicatedLaneRoute(lane, route) { + if lane != nil { + retiredSidecar = lane.sidecar + retiredFlight = lane.attachFlight + } + lane = &bulkDedicatedLane{id: laneID, route: route} c.bulkDedicatedLanes[laneID] = lane } if lane.sidecar != nil { - return lane.sidecar, nil, false + activeSidecar := lane.sidecar + c.bulkDedicatedSidecarMu.Unlock() + if retiredSidecar != nil { + retiredSidecar.close() + } + if retiredFlight != nil { + retiredFlight.finish(errTransportDetached) + } + return activeSidecar, nil, false, nil } if lane.attachFlight != nil { - return nil, lane.attachFlight, false + pendingFlight := lane.attachFlight + c.bulkDedicatedSidecarMu.Unlock() + if retiredSidecar != nil { + retiredSidecar.close() + } + if retiredFlight != nil { + retiredFlight.finish(errTransportDetached) + } + return nil, pendingFlight, false, nil } flight := newBulkDedicatedAttachFlight() lane.attachFlight = flight - return nil, flight, true + c.bulkDedicatedSidecarMu.Unlock() + if retiredSidecar != nil { + retiredSidecar.close() + } + if retiredFlight != nil { + retiredFlight.finish(errTransportDetached) + } + return nil, flight, true, nil } func (c *ClientCommon) finishClientDedicatedSidecarAttach(laneID uint32, flight *bulkDedicatedAttachFlight, err error) { @@ -245,25 +397,59 @@ func (c *ClientCommon) finishClientDedicatedSidecarAttach(laneID uint32, flight } func (c *ClientCommon) installClientDedicatedSidecar(laneID uint32, sidecar *bulkDedicatedSidecar) (*bulkDedicatedSidecar, bool) { + active, installed, _ := c.installClientDedicatedSidecarAtRoute(laneID, sidecar, c.clientSessionRouteSnapshot()) + return active, installed +} + +func (c *ClientCommon) installClientDedicatedSidecarAtRoute(laneID uint32, sidecar *bulkDedicatedSidecar, route clientSessionRoute) (*bulkDedicatedSidecar, bool, error) { if c == nil || sidecar == nil { - return nil, false + return nil, false, errBulkClientNil + } + if route.bound() { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return nil, false, err + } } laneID = normalizeBulkDedicatedLaneID(laneID) c.bulkDedicatedSidecarMu.Lock() - defer c.bulkDedicatedSidecarMu.Unlock() + if route.bound() && !c.clientSessionRouteCurrent(route) { + c.bulkDedicatedSidecarMu.Unlock() + return nil, false, transportDetachedSessionEpochError() + } if c.bulkDedicatedLanes == nil { c.bulkDedicatedLanes = make(map[uint32]*bulkDedicatedLane) } lane := c.bulkDedicatedLanes[laneID] - if lane == nil { - lane = &bulkDedicatedLane{id: laneID} + var retiredSidecar *bulkDedicatedSidecar + var retiredFlight *bulkDedicatedAttachFlight + if lane == nil || !sameClientDedicatedLaneRoute(lane, route) { + if lane != nil { + retiredSidecar = lane.sidecar + retiredFlight = lane.attachFlight + } + lane = &bulkDedicatedLane{id: laneID, route: route} c.bulkDedicatedLanes[laneID] = lane } if lane.sidecar != nil { - return lane.sidecar, false + activeSidecar := lane.sidecar + c.bulkDedicatedSidecarMu.Unlock() + if retiredSidecar != nil { + retiredSidecar.close() + } + if retiredFlight != nil { + retiredFlight.finish(errTransportDetached) + } + return activeSidecar, false, nil } lane.sidecar = sidecar - return sidecar, true + c.bulkDedicatedSidecarMu.Unlock() + if retiredSidecar != nil { + retiredSidecar.close() + } + if retiredFlight != nil { + retiredFlight.finish(errTransportDetached) + } + return sidecar, true, nil } func (c *ClientCommon) clearClientDedicatedSidecar(laneID uint32, sidecar *bulkDedicatedSidecar) bool { @@ -285,9 +471,16 @@ func (c *ClientCommon) clearClientDedicatedSidecar(laneID uint32, sidecar *bulkD } func (c *ClientCommon) closeClientDedicatedSidecar() { + c.closeClientDedicatedSidecarWithError(errServiceShutdown) +} + +func (c *ClientCommon) closeClientDedicatedSidecarWithError(closeErr error) { if c == nil { return } + if closeErr == nil { + closeErr = errServiceShutdown + } c.bulkDedicatedSidecarMu.Lock() lanes := c.bulkDedicatedLanes c.bulkDedicatedLanes = make(map[uint32]*bulkDedicatedLane) @@ -300,7 +493,7 @@ func (c *ClientCommon) closeClientDedicatedSidecar() { lane.sidecar.close() } if lane.attachFlight != nil { - lane.attachFlight.finish(errServiceShutdown) + lane.attachFlight.finish(closeErr) } } } @@ -358,6 +551,20 @@ func (s *ServerCommon) serverDedicatedSidecarSnapshotForLane(logical *LogicalCon return nil } +func (s *ServerCommon) serverDedicatedSidecarCurrent(logical *LogicalConn, sidecar *bulkDedicatedSidecar) bool { + if s == nil || logical == nil || sidecar == nil { + return false + } + return s.serverDedicatedSidecarSnapshotForLane(logical, sidecar.laneID) == sidecar +} + +func bulkDedicatedSidecarConnCurrent(bulk *bulkHandle, sidecar *bulkDedicatedSidecar) bool { + if bulk == nil || sidecar == nil || sidecar.conn == nil { + return false + } + return bulk.dedicatedConnSnapshot() == sidecar.conn +} + func (s *ServerCommon) installServerDedicatedSidecar(logical *LogicalConn, laneID uint32, sidecar *bulkDedicatedSidecar) *bulkDedicatedSidecar { if s == nil || logical == nil || sidecar == nil { return nil @@ -447,5 +654,7 @@ func (s *ServerCommon) attachServerDedicatedSidecarIfExists(logical *LogicalConn if sidecar == nil || sidecar.conn == nil { return } - _ = bulk.attachDedicatedConnShared(sidecar.conn) + _ = sidecar.withConn(func(conn net.Conn) error { + return bulk.attachDedicatedConnShared(conn) + }) } diff --git a/bulk_dispatcher.go b/bulk_dispatcher.go index d899d57..a8489b3 100644 --- a/bulk_dispatcher.go +++ b/bulk_dispatcher.go @@ -12,32 +12,43 @@ import ( const bulkDispatchRejectTimeout = 300 * time.Millisecond func (c *ClientCommon) dispatchFastBulkFrame(frame bulkFastFrame) { - c.dispatchFastBulkFrameWithOwner(frame, nil) + c.dispatchFastBulkFrameAtRoute(c.clientSessionRouteSnapshot(), frame) } func (c *ClientCommon) dispatchFastBulkFrameWithOwner(frame bulkFastFrame, owner *bulkReadPayloadOwner) { + c.dispatchFastBulkFrameWithOwnerAtRoute(c.clientSessionRouteSnapshot(), frame, owner) +} + +func (c *ClientCommon) dispatchFastBulkFrameAtRoute(route clientSessionRoute, frame bulkFastFrame) { + c.dispatchFastBulkFrameWithOwnerAtRoute(route, frame, nil) +} + +func (c *ClientCommon) dispatchFastBulkFrameWithOwnerAtRoute(route clientSessionRoute, frame bulkFastFrame, owner *bulkReadPayloadOwner) { if frame.DataID == 0 { return } + if route.bound() && !c.clientSessionRouteCurrent(route) { + return + } runtime := c.getBulkRuntime() if runtime == nil { return } - bulk, ok := runtime.lookupByDataID(clientFileScope(), frame.DataID) + bulk, ok := runtime.lookupInboundFrame(clientFileScope(), frame.DataID) if !ok { if c.showError || c.debugMode { fmt.Println("client bulk data for unknown data id", frame.DataID) } - c.bestEffortRejectInboundBulkData("", frame.DataID, errBulkNotFound.Error()) + c.bestEffortRejectInboundBulkDataAtRoute(route, "", frame.DataID, errBulkNotFound.Error()) return } - if !bulk.acceptsClientSessionEpoch(c.currentClientSessionEpoch()) { + if !bulk.acceptsClientSessionRoute(route) { if c.showError || c.debugMode { fmt.Println("client bulk data rejected by stale session epoch", frame.DataID) } detachErr := transportDetachedSessionEpochError() bulk.markReset(detachErr) - c.bestEffortRejectInboundBulkData(bulk.ID(), frame.DataID, detachErr.Error()) + c.bestEffortRejectInboundBulkDataAtRoute(route, bulk.ID(), frame.DataID, detachErr.Error()) return } switch frame.Type { @@ -53,7 +64,7 @@ func (c *ClientCommon) dispatchFastBulkFrameWithOwner(frame bulkFastFrame, owner fmt.Println("client bulk push chunk error", err) } if !errors.Is(err, io.EOF) { - c.bestEffortRejectInboundBulkData(bulk.ID(), frame.DataID, err.Error()) + c.bestEffortRejectInboundBulkDataAtRoute(route, bulk.ID(), frame.DataID, err.Error()) } } case bulkFastPayloadTypeClose: @@ -74,7 +85,7 @@ func (c *ClientCommon) dispatchFastBulkFrameWithOwner(frame bulkFastFrame, owner if c.showError || c.debugMode { fmt.Println("client bulk release decode error", err) } - c.bestEffortRejectInboundBulkData(bulk.ID(), frame.DataID, err.Error()) + c.bestEffortRejectInboundBulkDataAtRoute(route, bulk.ID(), frame.DataID, err.Error()) return } bulk.releaseOutboundWindow(bytes, chunks) @@ -97,7 +108,7 @@ func (s *ServerCommon) dispatchFastBulkFrameWithOwner(logical *LogicalConn, tran if runtime == nil { return } - bulk, ok := runtime.lookupByDataID(serverFileScope(logical), frame.DataID) + bulk, ok := runtime.lookupInboundFrame(serverFileScope(logical), frame.DataID) if !ok { if s.showError || s.debugMode { fmt.Println("server bulk data for unknown data id", frame.DataID) @@ -221,12 +232,16 @@ func (s *ServerCommon) tryDispatchBorrowedBulkTransportPayload(source interface{ } func (c *ClientCommon) bestEffortRejectInboundBulkData(bulkID string, dataID uint64, message string) { + c.bestEffortRejectInboundBulkDataAtRoute(c.clientSessionRouteSnapshot(), bulkID, dataID, message) +} + +func (c *ClientCommon) bestEffortRejectInboundBulkDataAtRoute(route clientSessionRoute, bulkID string, dataID uint64, message string) { if c == nil || (bulkID == "" && dataID == 0) { return } ctx, cancel := context.WithTimeout(context.Background(), bulkDispatchRejectTimeout) defer cancel() - _, _ = sendBulkResetClient(ctx, c, BulkResetRequest{ + _, _ = sendBulkResetClientAtRoute(ctx, c, route, BulkResetRequest{ BulkID: bulkID, DataID: dataID, Error: message, diff --git a/bulk_fastpath.go b/bulk_fastpath.go index d635e2d..035407d 100644 --- a/bulk_fastpath.go +++ b/bulk_fastpath.go @@ -41,6 +41,9 @@ func encodeBulkFastFrameHeader(dst []byte, frameType uint8, flags uint8, dataID if dataID == 0 { return errBulkDataIDEmpty } + if payloadLen < 0 || uint64(payloadLen) > uint64(^uint32(0)) { + return errBulkFastPayloadInvalid + } if len(dst) < bulkFastPayloadHeaderLen { return errBulkFastPayloadInvalid } @@ -92,8 +95,8 @@ func decodeBulkFastFrame(payload []byte) (bulkFastFrame, bool, error) { default: return bulkFastFrame{}, true, errBulkFastPayloadInvalid } - dataLen := int(binary.BigEndian.Uint32(payload[24:28])) - if dataLen < 0 || len(payload) != bulkFastPayloadHeaderLen+dataLen { + wireDataLen := binary.BigEndian.Uint32(payload[24:28]) + if uint64(len(payload)-bulkFastPayloadHeaderLen) != uint64(wireDataLen) { return bulkFastFrame{}, true, errBulkFastPayloadInvalid } dataID := binary.BigEndian.Uint64(payload[8:16]) @@ -200,10 +203,14 @@ func (c *ClientCommon) encodeBulkFastBatchPayloadPooled(frames []bulkFastFrame) } func (c *ClientCommon) sendFastBulkData(ctx context.Context, dataID uint64, seq uint64, chunk []byte, fastPathVersion uint8) error { - binding := c.clientTransportBindingSnapshot() - if binding == nil { - return net.ErrClosed + return c.sendFastBulkDataAtRoute(ctx, c.clientSessionRouteSnapshot(), dataID, seq, chunk, fastPathVersion) +} + +func (c *ClientCommon) sendFastBulkDataAtRoute(ctx context.Context, route clientSessionRoute, dataID uint64, seq uint64, chunk []byte, fastPathVersion uint8) error { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return err } + binding := route.binding if sender := binding.clientBulkBatchSenderSnapshot(c); sender != nil { return sender.submitData(ctx, dataID, seq, fastPathVersion, chunk) } @@ -211,17 +218,21 @@ func (c *ClientCommon) sendFastBulkData(ctx context.Context, dataID uint64, seq if err != nil { return err } - return c.writePayloadToTransport(payload) + return c.writePayloadToTransportBindingContextTimeout(ctx, binding, payload, 0) } func (c *ClientCommon) sendFastBulkWrite(ctx context.Context, dataID uint64, startSeq uint64, chunkSize int, fastPathVersion uint8, payload []byte, payloadOwned bool) (int, error) { + return c.sendFastBulkWriteAtRoute(ctx, c.clientSessionRouteSnapshot(), dataID, startSeq, chunkSize, fastPathVersion, payload, payloadOwned) +} + +func (c *ClientCommon) sendFastBulkWriteAtRoute(ctx context.Context, route clientSessionRoute, dataID uint64, startSeq uint64, chunkSize int, fastPathVersion uint8, payload []byte, payloadOwned bool) (int, error) { if len(payload) == 0 { return 0, nil } - binding := c.clientTransportBindingSnapshot() - if binding == nil { - return 0, net.ErrClosed + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return 0, err } + binding := route.binding if sender := binding.clientBulkBatchSenderSnapshot(c); sender != nil { return sender.submitWrite(ctx, dataID, startSeq, fastPathVersion, payload, chunkSize, payloadOwned) } @@ -235,7 +246,7 @@ func (c *ClientCommon) sendFastBulkWrite(ctx context.Context, dataID uint64, sta if end > len(payload) { end = len(payload) } - if err := c.sendFastBulkData(ctx, dataID, seq, payload[written:end], fastPathVersion); err != nil { + if err := c.sendFastBulkDataAtRoute(ctx, route, dataID, seq, payload[written:end], fastPathVersion); err != nil { return written, err } seq++ @@ -245,6 +256,10 @@ func (c *ClientCommon) sendFastBulkWrite(ctx context.Context, dataID uint64, sta } func (c *ClientCommon) sendFastBulkControl(ctx context.Context, frameType uint8, flags uint8, dataID uint64, seq uint64, fastPathVersion uint8, payload []byte) error { + return c.sendFastBulkControlAtRoute(ctx, c.clientSessionRouteSnapshot(), frameType, flags, dataID, seq, fastPathVersion, payload) +} + +func (c *ClientCommon) sendFastBulkControlAtRoute(ctx context.Context, route clientSessionRoute, frameType uint8, flags uint8, dataID uint64, seq uint64, fastPathVersion uint8, payload []byte) error { frame := bulkFastFrame{ Type: frameType, Flags: flags, @@ -252,10 +267,10 @@ func (c *ClientCommon) sendFastBulkControl(ctx context.Context, frameType uint8, Seq: seq, Payload: payload, } - binding := c.clientTransportBindingSnapshot() - if binding == nil { - return net.ErrClosed + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return err } + binding := route.binding if sender := binding.clientBulkBatchSenderSnapshot(c); sender != nil { return sender.submitControl(ctx, frameType, flags, dataID, seq, fastPathVersion, payload) } @@ -263,7 +278,7 @@ func (c *ClientCommon) sendFastBulkControl(ctx context.Context, frameType uint8, if err != nil { return err } - return c.writePayloadToTransport(encoded) + return c.writePayloadToTransportBindingContextTimeout(ctx, binding, encoded, 0) } func (c *ClientCommon) encodeBulkFastControlPayload(frameType uint8, flags uint8, dataID uint64, seq uint64, payload []byte) ([]byte, error) { @@ -359,7 +374,7 @@ func (s *ServerCommon) sendFastBulkDataTransport(ctx context.Context, logical *L if logical == nil { return errTransportDetached } - if binding := logical.transportBindingSnapshot(); binding != nil { + if binding := serverTransportBindingSnapshot(logical, transport); binding != nil { if binding.queueSnapshot() != nil { if sender := binding.serverBulkBatchSenderSnapshot(logical); sender != nil { return sender.submitData(ctx, dataID, seq, fastPathVersion, chunk) @@ -386,7 +401,7 @@ func (s *ServerCommon) sendFastBulkWriteTransport(ctx context.Context, logical * if logical == nil { return 0, errTransportDetached } - if binding := logical.transportBindingSnapshot(); binding != nil { + if binding := serverTransportBindingSnapshot(logical, transport); binding != nil { if binding.queueSnapshot() != nil { if sender := binding.serverBulkBatchSenderSnapshot(logical); sender != nil { return sender.submitWrite(ctx, dataID, startSeq, fastPathVersion, payload, chunkSize, payloadOwned) @@ -422,7 +437,7 @@ func (s *ServerCommon) sendFastBulkControlTransport(ctx context.Context, logical if logical == nil { return errTransportDetached } - if binding := logical.transportBindingSnapshot(); binding != nil { + if binding := serverTransportBindingSnapshot(logical, transport); binding != nil { if binding.queueSnapshot() != nil { if sender := binding.serverBulkBatchSenderSnapshot(logical); sender != nil { return sender.submitControl(ctx, frameType, flags, dataID, seq, fastPathVersion, payload) @@ -517,11 +532,15 @@ func decryptTransportPayloadWithFallbackPooled(primary transportProtectionProfil } func (c *ClientCommon) tryDispatchBorrowedTransportPlain(plain []byte, release func()) bool { + return c.tryDispatchBorrowedTransportPlainAtRoute(c.clientSessionRouteSnapshot(), plain, release) +} + +func (c *ClientCommon) tryDispatchBorrowedTransportPlainAtRoute(route clientSessionRoute, plain []byte, release func()) bool { switch transportFastPayloadMagic(plain) { case bulkFastPayloadMagic, bulkFastBatchMagic: owner := newBulkReadPayloadOwner(release) matched, walkErr := walkBulkFastFrames(plain, func(frame bulkFastFrame) error { - c.dispatchFastBulkFrameWithOwner(frame, owner) + c.dispatchFastBulkFrameWithOwnerAtRoute(route, frame, owner) return nil }) if owner != nil { @@ -537,7 +556,7 @@ func (c *ClientCommon) tryDispatchBorrowedTransportPlain(plain []byte, release f case streamFastPayloadMagic, streamFastBatchMagic: owner := newStreamReadPayloadOwner(release) matched, walkErr := walkStreamFastFrames(plain, func(frame streamFastDataFrame) error { - c.dispatchFastStreamDataWithOwner(frame, owner) + c.dispatchFastStreamDataWithOwnerAtRoute(route, frame, owner) return nil }) if owner != nil { @@ -595,22 +614,33 @@ func (s *ServerCommon) tryDispatchBorrowedTransportPlain(logical *LogicalConn, t } func (c *ClientCommon) dispatchInboundTransportPayload(payload []byte, now time.Time) error { + return c.dispatchInboundTransportPayloadAtRoute(c.clientSessionRouteSnapshot(), payload, now) +} + +func (c *ClientCommon) dispatchInboundTransportPayloadAtRoute(route clientSessionRoute, payload []byte, now time.Time) error { plain, err := c.decryptTransportPayload(payload) if err != nil { return err } - return c.dispatchInboundTransportPlain(plain, now) + return c.dispatchInboundTransportPlainAtRoute(route, plain, now) } func (c *ClientCommon) dispatchInboundTransportPlain(plain []byte, now time.Time) error { + return c.dispatchInboundTransportPlainAtRoute(c.clientSessionRouteSnapshot(), plain, now) +} + +func (c *ClientCommon) dispatchInboundTransportPlainAtRoute(route clientSessionRoute, plain []byte, now time.Time) error { + if route.bound() && !c.clientSessionRouteCurrent(route) { + return transportDetachedSessionEpochError() + } if matched, err := walkBulkFastFrames(plain, func(frame bulkFastFrame) error { - c.dispatchFastBulkFrame(frame) + c.dispatchFastBulkFrameAtRoute(route, frame) return nil }); matched { return err } if matched, err := walkStreamFastFrames(plain, func(frame streamFastDataFrame) error { - c.dispatchFastStreamData(frame) + c.dispatchFastStreamDataWithOwnerAtRoute(route, frame, nil) return nil }); matched { return err @@ -619,7 +649,7 @@ func (c *ClientCommon) dispatchInboundTransportPlain(plain []byte, now time.Time if err != nil { return err } - c.dispatchEnvelope(env, now) + c.dispatchEnvelopeAtRoute(route, env, now) return nil } diff --git a/bulk_lifecycle_lease_test.go b/bulk_lifecycle_lease_test.go new file mode 100644 index 0000000..1df0527 --- /dev/null +++ b/bulk_lifecycle_lease_test.go @@ -0,0 +1,183 @@ +package notify + +import ( + "context" + "errors" + "io" + "testing" + "time" +) + +func TestBulkRuntimeAdoptFailureFinalizesCandidate(t *testing.T) { + tests := []struct { + name string + adopt func(*bulkRuntime, string, *bulkHandle) error + }{ + { + name: "inbound", + adopt: func(runtime *bulkRuntime, scope string, bulk *bulkHandle) error { + return runtime.adoptInbound(scope, bulk) + }, + }, + { + name: "reserved", + adopt: func(runtime *bulkRuntime, scope string, bulk *bulkHandle) error { + return runtime.adoptReserved(scope, bulk) + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + runtime := newBulkRuntime("cblk") + scope := clientFileScope() + existing := newBulkHandle(context.Background(), runtime, scope, BulkOpenRequest{ + BulkID: "duplicate", + DataID: 1, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if err := runtime.registerInbound(scope, existing); err != nil { + t.Fatalf("register existing bulk: %v", err) + } + defer existing.markReset(io.ErrClosedPipe) + + candidate := newBulkHandle(context.Background(), runtime, scope, BulkOpenRequest{ + BulkID: "duplicate", + DataID: 3, + ChunkSize: 4, + WindowBytes: 4, + MaxInFlight: 1, + }, 0, nil, nil, 0, nil, nil, nil, + func(context.Context, *bulkHandle, uint64, []byte, bool) (int, error) { + return 0, nil + }, + func(*bulkHandle, int64, int) error { return nil }, + ) + if err := test.adopt(runtime, scope, candidate); !errors.Is(err, errBulkAlreadyExists) { + t.Fatalf("adopt error = %v, want %v", err, errBulkAlreadyExists) + } + if err := candidate.resetErrSnapshot(); !errors.Is(err, errBulkAlreadyExists) { + t.Fatalf("candidate reset error = %v, want %v", err, errBulkAlreadyExists) + } + waitBulkWorkerStopped(t, "write", candidate.writeWorkerDone) + waitBulkWorkerStopped(t, "release", candidate.releaseWorkerDone) + }) + } +} + +func TestSharedBulkFinalizeDoesNotReleaseDedicatedLane(t *testing.T) { + client := NewClient().(*ClientCommon) + laneID := client.reserveBulkDedicatedLane() + defer client.releaseBulkDedicatedLane(laneID) + + bulk := newBulkHandle(context.Background(), nil, clientFileScope(), BulkOpenRequest{ + BulkID: "shared", + DataID: 1, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.setClientSnapshotOwner(client) + bulk.finalize() + + if got := clientDedicatedLaneActiveBulks(client, laneID); got != 1 { + t.Fatalf("shared finalize changed lane %d active bulks to %d, want 1", laneID, got) + } +} + +func TestDedicatedLaneLeaseRetainsExactLaneAndReleasesOnce(t *testing.T) { + client := NewClient().(*ClientCommon) + laneID := client.retainBulkDedicatedLane(1) + bulk := newBulkHandle(context.Background(), nil, clientFileScope(), BulkOpenRequest{ + BulkID: "dedicated", + DataID: 1, + Dedicated: true, + DedicatedLaneID: laneID, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.setClientSnapshotOwner(client) + bulk.markDedicatedLaneReserved() + + otherLaneID := client.reserveBulkDedicatedLane() + if otherLaneID == laneID { + t.Fatalf("new lane reservation overwrote retained lane %d", laneID) + } + defer client.releaseBulkDedicatedLane(otherLaneID) + + bulk.finalize() + bulk.finalize() + if got := clientDedicatedLaneActiveBulks(client, laneID); got != 0 { + t.Fatalf("dedicated lane %d active bulks after duplicate finalize = %d, want 0", laneID, got) + } +} + +func TestInboundDedicatedRetainFailureFinalizesCandidate(t *testing.T) { + client := NewClient().(*ClientCommon) + runtime := newBulkRuntime("retain-failure") + route := clientSessionRoute{epoch: 1} + bulk := newBulkHandle(context.Background(), runtime, clientFileScope(), BulkOpenRequest{ + BulkID: "retain-failure", + DataID: 1, + Dedicated: true, + DedicatedLaneID: 9, + ChunkSize: 4, + WindowBytes: 4, + MaxInFlight: 1, + }, 1, nil, nil, 0, nil, nil, nil, + func(context.Context, *bulkHandle, uint64, []byte, bool) (int, error) { + return 0, nil + }, + func(*bulkHandle, int64, int) error { return nil }, + ) + bulk.setClientSnapshotOwner(client) + bulk.setClientSessionRoute(route) + + err := client.retainBulkDedicatedLaneAtRoute(bulk.dedicatedLaneIDSnapshot(), route) + if err == nil { + t.Fatal("retain should fail for an unavailable session route") + } + bulk.markReset(err) + + if got := bulk.resetErrSnapshot(); got == nil { + t.Fatal("retain failure did not set candidate reset error") + } + if _, ok := runtime.lookup(clientFileScope(), bulk.ID()); ok { + t.Fatal("unadopted candidate was registered in bulk runtime") + } + waitBulkWorkerStopped(t, "write", bulk.writeWorkerDone) + waitBulkWorkerStopped(t, "release", bulk.releaseWorkerDone) +} + +func TestFinalizedBulkRejectsDedicatedSenderInstall(t *testing.T) { + bulk := newBulkHandle(context.Background(), nil, clientFileScope(), BulkOpenRequest{ + BulkID: "finalized-sender", + DataID: 1, + Dedicated: true, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.finalize() + + sender := &bulkDedicatedSender{} + if got := bulk.installDedicatedSender(sender); got != nil { + t.Fatalf("install sender after finalize = %p, want nil", got) + } + if got := bulk.dedicatedSenderSnapshot(); got != nil { + t.Fatalf("finalized bulk retained sender %p", got) + } +} + +func waitBulkWorkerStopped(t *testing.T, name string, done <-chan struct{}) { + t.Helper() + if done == nil { + t.Fatalf("%s worker was not started", name) + } + select { + case <-done: + case <-time.After(time.Second): + t.Fatalf("%s worker did not stop", name) + } +} + +func clientDedicatedLaneActiveBulks(client *ClientCommon, laneID uint32) int { + client.bulkDedicatedSidecarMu.Lock() + defer client.bulkDedicatedSidecarMu.Unlock() + lane := client.bulkDedicatedLanes[normalizeBulkDedicatedLaneID(laneID)] + if lane == nil { + return 0 + } + return lane.activeBulks +} diff --git a/bulk_recovery.go b/bulk_recovery.go new file mode 100644 index 0000000..7073e1a --- /dev/null +++ b/bulk_recovery.go @@ -0,0 +1,279 @@ +package notify + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "sync" + "time" +) + +const ( + bulkRecoveryQueueSize = 64 + bulkRecoveryWorkers = 4 + bulkRecoveryAttempts = 3 + bulkRecoveryAttempt = 750 * time.Millisecond + bulkRecoveryBackoff = 50 * time.Millisecond +) + +var errBulkRecoveryQueueFull = errors.New("bulk recovery queue full") + +type bulkRecoveryTask func(context.Context) error + +// bulkRecoveryQueue bounds cleanup work after failed bulk opens. A small fixed +// worker set prevents both unbounded goroutines and multi-minute serial drain. +type bulkRecoveryQueue struct { + mu sync.Mutex + tasks []bulkRecoveryTask + workerRunning bool + workersRunning int + onError func(error) +} + +func newBulkRecoveryQueue(onError func(error)) *bulkRecoveryQueue { + return &bulkRecoveryQueue{ + onError: onError, + } +} + +func (q *bulkRecoveryQueue) enqueue(task bulkRecoveryTask) bool { + if task == nil { + return true + } + if q == nil { + return false + } + q.mu.Lock() + if len(q.tasks) >= bulkRecoveryQueueSize { + q.mu.Unlock() + return false + } + q.tasks = append(q.tasks, task) + start := 0 + for q.workersRunning < bulkRecoveryWorkers && q.workersRunning < len(q.tasks) { + q.workersRunning++ + start++ + } + q.workerRunning = q.workersRunning > 0 + q.mu.Unlock() + for i := 0; i < start; i++ { + go q.loop() + } + return true +} + +func (q *bulkRecoveryQueue) loop() { + for { + q.mu.Lock() + if len(q.tasks) == 0 { + q.workersRunning-- + q.workerRunning = q.workersRunning > 0 + q.mu.Unlock() + return + } + task := q.tasks[0] + copy(q.tasks, q.tasks[1:]) + q.tasks[len(q.tasks)-1] = nil + q.tasks = q.tasks[:len(q.tasks)-1] + q.mu.Unlock() + q.execute(task) + } +} + +func (q *bulkRecoveryQueue) execute(task bulkRecoveryTask) { + if q == nil || task == nil { + return + } + q.run(task) +} + +func (q *bulkRecoveryQueue) run(task bulkRecoveryTask) { + if q == nil || task == nil { + return + } + lastErr := runBulkRecoveryTaskContext(context.Background(), task) + if lastErr != nil && q.onError != nil { + q.onError(lastErr) + } +} + +// runBulkRecoveryTask performs a bounded reset attempt synchronously. Callers +// that must establish ordering (for example, Auto dedicated -> shared fallback) +// use this path so a later open cannot overtake cleanup queued in the background. +func runBulkRecoveryTask(task bulkRecoveryTask) error { + return runBulkRecoveryTaskContext(context.Background(), task) +} + +func runBulkRecoveryTaskContext(parent context.Context, task bulkRecoveryTask) error { + if task == nil { + return nil + } + if parent == nil { + parent = context.Background() + } + deadline := time.Now().Add(bulkOpenRecoveryTimeout) + var lastErr error + for attempt := 0; attempt < bulkRecoveryAttempts; attempt++ { + if err := parent.Err(); err != nil { + return err + } + remaining := time.Until(deadline) + if remaining <= 0 { + break + } + attemptTimeout := remaining + if attemptTimeout > bulkRecoveryAttempt { + attemptTimeout = bulkRecoveryAttempt + } + ctx, cancel := context.WithTimeout(parent, attemptTimeout) + err := task(ctx) + cancel() + if err == nil { + return nil + } + lastErr = err + if !bulkRecoveryErrorRetryable(err) { + break + } + if attempt+1 >= bulkRecoveryAttempts { + break + } + backoff := bulkRecoveryBackoff << attempt + if backoff > time.Until(deadline) { + backoff = time.Until(deadline) + } + if backoff > 0 { + timer := time.NewTimer(backoff) + select { + case <-parent.Done(): + if !timer.Stop() { + <-timer.C + } + return parent.Err() + case <-timer.C: + } + } + } + return lastErr +} + +func bulkRecoveryErrorRetryable(err error) bool { + if err == nil { + return false + } + return !errors.Is(err, errTransportDetached) && + !errors.Is(err, errServiceShutdown) && + !errors.Is(err, net.ErrClosed) && + !errors.Is(err, io.ErrClosedPipe) +} + +func (c *ClientCommon) resetBulkAtRouteAndWait(ctx context.Context, route clientSessionRoute, req BulkResetRequest) error { + if c == nil { + return errBulkClientNil + } + err := runBulkRecoveryTaskContext(ctx, newClientBulkResetRecoveryTaskAtRoute(c, route, req)) + if err != nil && bulkRecoveryErrorRetryable(err) { + c.bestEffortBulkResetAtRoute(route, req) + } + return err +} + +func (c *ClientCommon) cleanupBulkResetAtRoute(ctx context.Context, route clientSessionRoute, req BulkResetRequest, wait bool) error { + if wait { + return c.resetBulkAtRouteAndWait(ctx, route, req) + } + c.bestEffortBulkResetAtRoute(route, req) + return nil +} + +func (s *ServerCommon) resetBulkLogicalAndWait(ctx context.Context, logical *LogicalConn, transport *TransportConn, req BulkResetRequest) error { + if s == nil { + return errBulkServerNil + } + err := runBulkRecoveryTaskContext(ctx, newServerBulkResetRecoveryTask(s, logical, transport, req)) + if err != nil && bulkRecoveryErrorRetryable(err) { + s.bestEffortBulkResetLogical(logical, transport, req) + } + return err +} + +func (s *ServerCommon) cleanupBulkLogicalReset(ctx context.Context, logical *LogicalConn, transport *TransportConn, req BulkResetRequest, wait bool) error { + if wait { + return s.resetBulkLogicalAndWait(ctx, logical, transport, req) + } + s.bestEffortBulkResetLogical(logical, transport, req) + return nil +} + +func (s *ServerCommon) resetBulkTransportAndWait(ctx context.Context, transport *TransportConn, req BulkResetRequest) error { + if s == nil { + return errBulkServerNil + } + err := runBulkRecoveryTaskContext(ctx, newServerBulkResetRecoveryTask(s, transport.logicalConnSnapshot(), transport, req)) + if err != nil && bulkRecoveryErrorRetryable(err) { + s.bestEffortBulkResetTransport(transport, req) + } + return err +} + +func (s *ServerCommon) cleanupBulkTransportReset(ctx context.Context, transport *TransportConn, req BulkResetRequest, wait bool) error { + if wait { + return s.resetBulkTransportAndWait(ctx, transport, req) + } + s.bestEffortBulkResetTransport(transport, req) + return nil +} + +func (c *ClientCommon) handleBulkRecoveryOverflow(epoch uint64, req BulkResetRequest) { + route := c.clientSessionRouteSnapshot() + route.epoch = epoch + c.handleBulkRecoveryOverflowAtRoute(route, req) +} + +func (c *ClientCommon) handleBulkRecoveryOverflowAtRoute(route clientSessionRoute, req BulkResetRequest) { + if c == nil { + return + } + err := fmt.Errorf("%w: bulk=%s data=%d", errBulkRecoveryQueueFull, req.BulkID, req.DataID) + c.reportBulkRecoveryError(err) + if c.clientSessionRouteCurrent(route) && route.epoch != 0 { + c.stopClientSessionIfCurrent(route.epoch, "bulk recovery queue full", err) + } +} + +func (s *ServerCommon) handleBulkRecoveryOverflow(logical *LogicalConn, transport *TransportConn, req BulkResetRequest) { + if s == nil { + return + } + err := fmt.Errorf("%w: bulk=%s data=%d", errBulkRecoveryQueueFull, req.BulkID, req.DataID) + s.reportBulkRecoveryError(err) + if logical != nil && transport != nil && transport.IsCurrent() { + s.detachLogicalSessionTransport(logical, "bulk recovery queue full", err) + } +} + +func (c *ClientCommon) reportBulkRecoveryError(err error) { + if c == nil || err == nil { + return + } + c.mu.Lock() + debug := c.showError || c.debugMode + c.mu.Unlock() + if debug { + fmt.Printf("notify bulk reset recovery failed: %v\n", err) + } +} + +func (s *ServerCommon) reportBulkRecoveryError(err error) { + if s == nil || err == nil { + return + } + s.mu.RLock() + debug := s.showError || s.debugMode + s.mu.RUnlock() + if debug { + fmt.Printf("notify bulk reset recovery failed: %v\n", err) + } +} diff --git a/bulk_recovery_test.go b/bulk_recovery_test.go new file mode 100644 index 0000000..4bbe1f3 --- /dev/null +++ b/bulk_recovery_test.go @@ -0,0 +1,719 @@ +package notify + +import ( + "context" + "errors" + "math" + "net" + "os" + "sync/atomic" + "testing" + "time" + + "b612.me/stario" +) + +func TestBulkRecoveryQueueRetriesWithoutConcurrentWorkers(t *testing.T) { + var attempts atomic.Int32 + var active atomic.Int32 + var maxActive atomic.Int32 + done := make(chan struct{}) + q := newBulkRecoveryQueue(func(error) { t.Errorf("recovery should succeed") }) + if !q.enqueue(func(context.Context) error { + current := active.Add(1) + for { + previous := maxActive.Load() + if current <= previous || maxActive.CompareAndSwap(previous, current) { + break + } + } + defer active.Add(-1) + attempt := attempts.Add(1) + if attempt < 3 { + return errors.New("retry") + } + close(done) + return nil + }) { + t.Fatal("enqueue unexpectedly rejected task") + } + select { + case <-done: + case <-time.After(3 * time.Second): + t.Fatal("timed out waiting for recovery retries") + } + if got := attempts.Load(); got != 3 { + t.Fatalf("attempts = %d, want 3", got) + } + if got := maxActive.Load(); got != 1 { + t.Fatalf("max concurrent recovery workers = %d, want 1", got) + } +} + +func TestBulkRecoveryQueueStopsWorkerWhenIdleAndRestarts(t *testing.T) { + q := newBulkRecoveryQueue(func(error) { t.Fatal("unexpected recovery error") }) + done := make(chan struct{}) + if !q.enqueue(func(context.Context) error { + close(done) + return nil + }) { + t.Fatal("enqueue unexpectedly rejected task") + } + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("timed out waiting for recovery task") + } + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + q.mu.Lock() + running := q.workerRunning + q.mu.Unlock() + if !running { + break + } + time.Sleep(time.Millisecond) + } + q.mu.Lock() + running := q.workerRunning + q.mu.Unlock() + if running { + t.Fatal("recovery worker remained alive after queue drained") + } + + secondDone := make(chan struct{}) + if !q.enqueue(func(context.Context) error { + close(secondDone) + return nil + }) { + t.Fatal("enqueue after idle unexpectedly rejected task") + } + select { + case <-secondDone: + case <-time.After(time.Second): + t.Fatal("recovery queue did not restart after idle") + } +} + +func TestBulkRecoveryQueueDoesNotRetryDetachedSession(t *testing.T) { + var attempts atomic.Int32 + q := newBulkRecoveryQueue(nil) + q.run(func(context.Context) error { + attempts.Add(1) + return transportDetachedSessionEpochError() + }) + if got := attempts.Load(); got != 1 { + t.Fatalf("detached recovery attempts = %d, want 1", got) + } +} + +func TestBulkRecoveryContextCancellationStopsRetryLoop(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + started := make(chan struct{}) + var attempts atomic.Int32 + task := func(taskCtx context.Context) error { + attempts.Add(1) + close(started) + <-taskCtx.Done() + return taskCtx.Err() + } + done := make(chan error, 1) + go func() { done <- runBulkRecoveryTaskContext(ctx, task) }() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("recovery task did not start") + } + cancel() + select { + case err := <-done: + if !errors.Is(err, context.Canceled) { + t.Fatalf("recovery error = %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("recovery did not stop after parent cancellation") + } + if got := attempts.Load(); got != 1 { + t.Fatalf("recovery attempts = %d, want 1", got) + } +} + +func TestClientBulkRecoveryQueueOverflowDoesNotBlockCaller(t *testing.T) { + release := make(chan struct{}) + started := make(chan struct{}, bulkRecoveryQueueSize+16) + blockingTask := func(ctx context.Context) error { + started <- struct{}{} + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + } + q := newBulkRecoveryQueue(nil) + if !q.enqueue(blockingTask) { + t.Fatal("initial recovery task was rejected") + } + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("recovery worker did not start") + } + for q.enqueue(blockingTask) { + } + + client := NewClient().(*ClientCommon) + client.bulkRecovery = q + epoch := client.beginClientSessionEpoch() + returned := make(chan struct{}) + go func() { + client.bestEffortBulkResetAtEpoch(epoch, BulkResetRequest{BulkID: "overflow"}) + close(returned) + }() + select { + case <-returned: + close(release) + case <-time.After(100 * time.Millisecond): + close(release) + <-returned + t.Fatal("queue-full recovery blocked the caller") + } +} + +func TestClientBulkResetRecoveryTaskRejectsStaleSession(t *testing.T) { + client := NewClient().(*ClientCommon) + epoch := client.beginClientSessionEpoch() + client.beginClientSessionEpoch() + + task := newClientBulkResetRecoveryTask(client, epoch, BulkResetRequest{BulkID: "stale"}) + err := task(context.Background()) + if !errors.Is(err, errTransportDetached) { + t.Fatalf("stale recovery task error = %v, want transport detached", err) + } +} + +func TestBulkOpenAutoWaitsForDedicatedResetBeforeSharedFallback(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKServer failed: %v", err) + } + accepted := make(chan BulkAcceptInfo, 2) + server.SetBulkHandler(func(info BulkAcceptInfo) error { + accepted <- info + return nil + }) + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatalf("server Listen failed: %v", err) + } + defer server.Stop() + + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKClient failed: %v", err) + } + if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { + t.Fatalf("client Connect failed: %v", err) + } + defer client.Stop() + client.setClientConnectSource(newClientFactoryConnectSource(func(context.Context) (net.Conn, error) { + return nil, errors.New("forced attach dial failure") + })) + + // Occupy every asynchronous recovery worker. The old Auto path queued its + // reset behind these tasks and immediately sent shared open, reproducing the + // explicit-ID race. The fixed path performs the reset synchronously. + release := make(chan struct{}) + defer close(release) + started := make(chan struct{}, bulkRecoveryWorkers) + q := newBulkRecoveryQueue(nil) + blockingTask := func(ctx context.Context) error { + started <- struct{}{} + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + } + for i := 0; i < bulkRecoveryWorkers; i++ { + if !q.enqueue(blockingTask) { + t.Fatalf("enqueue blocking recovery task %d failed", i) + } + } + for i := 0; i < bulkRecoveryWorkers; i++ { + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("recovery worker did not become occupied") + } + } + client.bulkRecovery = q + + bulk, err := client.OpenBulk(context.Background(), BulkOpenOptions{ + Mode: BulkOpenModeAuto, + ID: "auto-reset-order", + Range: BulkRange{ + Offset: 0, + Length: 128, + }, + }) + if err != nil { + t.Fatalf("Auto fallback failed while reset workers were occupied: %v", err) + } + if bulk.Snapshot().Dedicated { + t.Fatal("Auto fallback returned a dedicated bulk") + } + defer bulk.Close() + + seenShared := false + deadline := time.After(2 * time.Second) + for !seenShared { + select { + case info := <-accepted: + if info.Bulk != nil { + defer info.Bulk.Close() + } + if info.ID == bulk.ID() { + seenShared = !info.Dedicated + } + case <-deadline: + t.Fatal("timed out waiting for shared fallback accept") + } + } +} + +func TestServerBulkResetRecoveryTaskRejectsStaleTransport(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) + defer runtimeCancel() + queue := stario.NewQueueCtx(runtimeCtx, 4, math.MaxUint32) + server.setServerSessionRuntime(&serverSessionRuntime{ + stopCtx: runtimeCtx, + stopFn: runtimeCancel, + queue: queue, + }) + server.markSessionStarted() + defer server.markSessionStopped("test done", nil) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical, _, _ := newRegisteredServerLogicalForTest(t, server, "bulk-recovery-stale", firstLeft, runtimeCtx, runtimeCancel) + firstTransport := logical.CurrentTransportConn() + if firstTransport == nil { + t.Fatal("first transport snapshot should exist") + } + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { + t.Fatalf("attachClientConnSessionTransport failed: %v", err) + } + if firstTransport.IsCurrent() { + t.Fatal("first transport should be stale after reattach") + } + + task := newServerBulkResetRecoveryTask(server, logical, firstTransport, BulkResetRequest{BulkID: "stale"}) + err := task(context.Background()) + if !errors.Is(err, errTransportDetached) { + t.Fatalf("stale recovery task error = %v, want transport detached", err) + } +} + +func TestClientBulkOpenDoesNotCrossTransportReattach(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + epoch := client.beginClientSessionEpoch() + client.setClientSessionRuntime(newClientSessionRuntime(firstLeft, stopCtx, stopFn, queue, epoch)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + + runtime := client.getBulkRuntime() + runtime.mu.Lock() + result := make(chan error, 1) + go func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err := client.OpenBulk(ctx, BulkOpenOptions{Mode: BulkOpenModeShared}) + result <- err + }() + waitForBulkOpenRouteCapture(t, runtime) + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := client.attachClientSessionTransport(secondLeft); err != nil { + runtime.mu.Unlock() + t.Fatalf("attach client replacement transport: %v", err) + } + runtime.mu.Unlock() + + if err := <-result; !errors.Is(err, errTransportDetached) { + t.Fatalf("bulk open error = %v, want transport detached", err) + } + assertNoPipeWrite(t, secondRight, "client bulk open crossed onto replacement transport") +} + +func TestServerBulkOpenDoesNotCrossTransportReattach(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) + defer runtimeCancel() + queue := stario.NewQueueCtx(runtimeCtx, 4, math.MaxUint32) + server.setServerSessionRuntime(&serverSessionRuntime{stopCtx: runtimeCtx, stopFn: runtimeCancel, queue: queue}) + server.markSessionStarted() + defer server.markSessionStopped("test done", nil) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical, _, _ := newRegisteredServerLogicalForTest(t, server, "bulk-open-reattach", firstLeft, runtimeCtx, runtimeCancel) + runtime := server.getBulkRuntime() + runtime.mu.Lock() + result := make(chan error, 1) + go func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + _, err := server.OpenBulkLogical(ctx, logical, BulkOpenOptions{Mode: BulkOpenModeShared}) + result <- err + }() + waitForBulkOpenRouteCapture(t, runtime) + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { + runtime.mu.Unlock() + t.Fatalf("attach server replacement transport: %v", err) + } + runtime.mu.Unlock() + + if err := <-result; !errors.Is(err, errTransportDetached) { + t.Fatalf("bulk open error = %v, want transport detached", err) + } + assertNoPipeWrite(t, secondRight, "server bulk open crossed onto replacement transport") +} + +func TestServerRejectsQueuedBulkOpenFromStaleTransport(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + var handlerCalls atomic.Int32 + server.SetBulkHandler(func(BulkAcceptInfo) error { + handlerCalls.Add(1) + return nil + }) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical := server.bootstrapAcceptedLogical("stale-inbound-bulk-open", nil, firstLeft) + if logical == nil { + t.Fatal("bootstrapAcceptedLogical should return logical") + } + staleTransport := logical.CurrentTransportConn() + if staleTransport == nil { + t.Fatal("initial transport snapshot should exist") + } + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { + t.Fatalf("attach replacement transport: %v", err) + } + payload, err := encode(BulkOpenRequest{BulkID: "queued-stale-open", DataID: 1}) + if err != nil { + t.Fatalf("encode BulkOpenRequest: %v", err) + } + message := Message{ + NetType: NET_SERVER, + LogicalConn: logical, + TransportConn: staleTransport, + TransferMsg: TransferMsg{ + Key: BulkOpenSignalKey, + Value: payload, + Type: MSG_ASYNC, + }, + } + + server.handleInboundBulkOpen(&message) + if got := handlerCalls.Load(); got != 0 { + t.Fatalf("stale BulkOpen handler calls = %d, want 0", got) + } + if bulk, ok := server.getBulkRuntime().lookup(serverFileScope(logical), "queued-stale-open"); ok { + t.Fatalf("stale BulkOpen registered runtime handle: %+v", bulk.snapshot()) + } +} + +func TestServerTransportReattachResetsExistingBulk(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical := server.bootstrapAcceptedLogical("bulk-reset-on-reattach", nil, firstLeft) + if logical == nil { + t.Fatal("bootstrapAcceptedLogical should return logical") + } + firstTransport := logical.CurrentTransportConn() + if firstTransport == nil { + t.Fatal("initial transport snapshot should exist") + } + runtime := server.getBulkRuntime() + bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, serverFileScope(logical), BulkOpenRequest{ + BulkID: "bulk-before-reattach", + DataID: 2, + }, 0, logical, firstTransport, firstTransport.TransportGeneration(), nil, nil, nil, nil, nil) + if err := runtime.register(serverFileScope(logical), bulk); err != nil { + t.Fatalf("register bulk: %v", err) + } + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := server.attachAcceptedLogicalTransport(logical, secondLeft.RemoteAddr(), secondLeft); err != nil { + t.Fatalf("attach replacement transport: %v", err) + } + + select { + case <-bulk.Context().Done(): + case <-time.After(time.Second): + t.Fatal("old-generation bulk remained active after transport reattach") + } + if err := bulk.resetErrSnapshot(); !errors.Is(err, errTransportDetached) { + t.Fatalf("old-generation bulk reset error = %v, want transport detached", err) + } + if registered, ok := runtime.lookup(serverFileScope(logical), bulk.ID()); ok { + t.Fatalf("old-generation bulk remained registered: %+v", registered.snapshot()) + } + if firstTransport.IsCurrent() { + t.Fatal("old transport remained current after replacement") + } +} + +func TestClientRejectsBulkOpenWhenRouteReattachesDuringRegistration(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + var handlerCalls atomic.Int32 + client.SetBulkHandler(func(BulkAcceptInfo) error { + handlerCalls.Add(1) + return nil + }) + + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + epoch := client.beginClientSessionEpoch() + client.setClientSessionRuntime(newClientSessionRuntime(firstLeft, stopCtx, stopFn, queue, epoch)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + + payload, err := encode(BulkOpenRequest{BulkID: "client-stale-inbound-open", DataID: 2}) + if err != nil { + t.Fatalf("encode BulkOpenRequest: %v", err) + } + message := Message{ + NetType: NET_CLIENT, + ServerConn: client, + clientRoute: client.clientSessionRouteSnapshot(), + TransferMsg: TransferMsg{ + Key: BulkOpenSignalKey, + Value: payload, + Type: MSG_ASYNC, + }, + } + + runtime := client.getBulkRuntime() + runtime.mu.Lock() + done := make(chan struct{}) + go func() { + defer close(done) + client.handleInboundBulkOpen(&message) + }() + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := client.attachClientSessionTransport(secondLeft); err != nil { + runtime.mu.Unlock() + t.Fatalf("attach client replacement transport: %v", err) + } + runtime.mu.Unlock() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("client stale BulkOpen handler did not return") + } + if got := handlerCalls.Load(); got != 0 { + t.Fatalf("stale client BulkOpen handler calls = %d, want 0", got) + } + if bulk, ok := runtime.lookup(clientFileScope(), "client-stale-inbound-open"); ok { + t.Fatalf("stale client BulkOpen registered runtime handle: %+v", bulk.snapshot()) + } +} + +func TestClientBulkReadyDoesNotCrossTransportReattach(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + epoch := client.beginClientSessionEpoch() + client.setClientSessionRuntime(newClientSessionRuntime(firstLeft, stopCtx, stopFn, queue, epoch)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + + route := client.clientSessionRouteSnapshot() + bulk := newBulkHandle(stopCtx, nil, clientFileScope(), BulkOpenRequest{ + BulkID: "client-ready-reattach", + DataID: 1, + }, epoch, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.setClientSessionRoute(route) + defer bulk.finalize() + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := client.attachClientSessionTransport(secondLeft); err != nil { + t.Fatalf("attach client replacement transport: %v", err) + } + + done := make(chan struct{}) + go func() { + client.clientBulkAcceptReadyNotifier(bulk)(nil) + close(done) + }() + assertNoPipeWrite(t, secondRight, "client bulk ready crossed onto replacement transport") + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("client bulk ready notifier did not reject stale route") + } +} + +func TestServerBulkReadyDoesNotCrossTransportReattach(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) + defer runtimeCancel() + queue := stario.NewQueueCtx(runtimeCtx, 4, math.MaxUint32) + server.setServerSessionRuntime(&serverSessionRuntime{stopCtx: runtimeCtx, stopFn: runtimeCancel, queue: queue}) + server.markSessionStarted() + defer server.markSessionStopped("test done", nil) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical, _, _ := newRegisteredServerLogicalForTest(t, server, "bulk-ready-reattach", firstLeft, runtimeCtx, runtimeCancel) + logical.applyClientConnAttachmentProfile(0, 100*time.Millisecond, server.defaultMsgEn, server.defaultMsgDe, server.handshakeRsaKey, server.SecretKey) + firstTransport := logical.CurrentTransportConn() + if firstTransport == nil { + t.Fatal("first transport snapshot should exist") + } + bulk := newBulkHandle(runtimeCtx, nil, serverFileScope(logical), BulkOpenRequest{ + BulkID: "server-ready-reattach", + DataID: 1, + }, 0, logical, firstTransport, firstTransport.TransportGeneration(), nil, nil, nil, nil, nil) + defer bulk.finalize() + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { + t.Fatalf("attach server replacement transport: %v", err) + } + + result := make(chan error, 1) + go func() { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + result <- sendBulkReadyServer(ctx, server, logical, firstTransport, BulkReadyRequest{ + BulkID: bulk.ID(), + DataID: bulk.dataIDSnapshot(), + }) + }() + assertNoPipeWrite(t, secondRight, "server bulk ready crossed onto replacement transport") + select { + case err := <-result: + if !errors.Is(err, errTransportDetached) { + t.Fatalf("server bulk ready error = %v, want transport detached", err) + } + case <-time.After(time.Second): + t.Fatal("server bulk ready send did not reject stale transport") + } +} + +func waitForBulkOpenRouteCapture(t *testing.T, runtime *bulkRuntime) { + t.Helper() + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + if runtime.seq.Load() != 0 { + return + } + time.Sleep(time.Millisecond) + } + runtime.mu.Unlock() + t.Fatal("timed out waiting for bulk open to capture its route") +} + +func assertNoPipeWrite(t *testing.T, conn net.Conn, message string) { + t.Helper() + if err := conn.SetReadDeadline(time.Now().Add(50 * time.Millisecond)); err != nil { + t.Fatalf("set pipe read deadline: %v", err) + } + buf := make([]byte, 1) + if n, err := conn.Read(buf); n != 0 || !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("%s: read=%d err=%v", message, n, err) + } +} + +func TestClientInboundParserDoesNotJoinFramesAcrossTransportReattach(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) + firstBinding := newTransportBinding(nil, queue) + secondBinding := newTransportBinding(nil, queue) + secondRuntime := prepareClientSessionRuntime(&clientSessionRuntime{ + transport: secondBinding, + transportAttached: true, + stopCtx: stopCtx, + stopFn: stopFn, + queue: queue, + inboundDispatcher: newInboundDispatcher(), + epoch: 2, + }) + client.setClientSessionRuntime(secondRuntime) + + received := make(chan Message, 2) + client.SetLink("reattach-frame", func(msg *Message) { received <- *msg }) + env, err := wrapTransferMsgEnvelope(TransferMsg{ID: 91, Key: "reattach-frame", Value: MsgVal("payload"), Type: MSG_ASYNC}, client.sequenceEn) + if err != nil { + t.Fatalf("wrap transfer envelope: %v", err) + } + wire, err := client.encodeEnvelope(env) + if err != nil { + t.Fatalf("encode envelope: %v", err) + } + cut := len(wire) / 2 + firstRoute := clientSessionRoute{binding: firstBinding, epoch: 1, sessionStopCtx: stopCtx, transportStopCtx: stopCtx} + secondRoute := clientSessionRouteFromRuntime(secondRuntime) + client.pushMessageFastAtRoute(firstRoute, queue, wire[:cut], secondRuntime.inboundDispatcher) + client.pushMessageFastAtRoute(secondRoute, queue, wire[cut:], secondRuntime.inboundDispatcher) + select { + case msg := <-received: + t.Fatalf("split frame crossed transports: %+v", msg.TransferMsg) + case <-time.After(20 * time.Millisecond): + } + + client.pushMessageFastAtRoute(secondRoute, queue, wire, secondRuntime.inboundDispatcher) + select { + case msg := <-received: + if msg.Key != "reattach-frame" || string(msg.Value) != "payload" { + t.Fatalf("decoded message = %+v", msg.TransferMsg) + } + case <-time.After(time.Second): + t.Fatal("complete replacement-transport frame was not dispatched") + } +} diff --git a/bulk_runtime.go b/bulk_runtime.go index 9ed8bed..8ed106d 100644 --- a/bulk_runtime.go +++ b/bulk_runtime.go @@ -11,19 +11,45 @@ import ( type bulkRuntime struct { rolePrefix string seq atomic.Uint64 - dataSeq atomic.Uint64 + dataSeq uint64 + dataStart uint64 + dataStep uint64 - mu sync.RWMutex - handler func(BulkAcceptInfo) error - bulks map[string]*bulkHandle - data map[string]map[uint64]*bulkHandle + mu sync.RWMutex + handler func(BulkAcceptInfo) error + bulks map[string]*bulkHandle + inbound map[string]map[uint64]*bulkHandle + outbound map[string]map[uint64]*bulkHandle + reserved map[string]map[uint64]struct{} } +type bulkDataIndexDirection uint8 + +const ( + bulkDataIndexInbound bulkDataIndexDirection = 1 << iota + bulkDataIndexOutbound + bulkDataIndexBoth = bulkDataIndexInbound | bulkDataIndexOutbound +) + func newBulkRuntime(rolePrefix string) *bulkRuntime { + dataStart, dataStep := uint64(1), uint64(1) + // Client- and server-originated IDs occupy disjoint wire namespaces. A + // bulk is duplex, so every incoming frame must be routable by DataID alone; + // partitioning the allocator prevents simultaneous opens from colliding. + if rolePrefix == "cblk" { + dataStep = 2 + } else if rolePrefix == "sblk" { + dataStart = 2 + dataStep = 2 + } return &bulkRuntime{ rolePrefix: rolePrefix, + dataStart: dataStart, + dataStep: dataStep, bulks: make(map[string]*bulkHandle), - data: make(map[string]map[uint64]*bulkHandle), + inbound: make(map[string]map[uint64]*bulkHandle), + outbound: make(map[string]map[uint64]*bulkHandle), + reserved: make(map[string]map[uint64]struct{}), } } @@ -38,7 +64,148 @@ func (r *bulkRuntime) nextDataID() uint64 { if r == nil { return 0 } - return r.dataSeq.Add(1) + r.mu.Lock() + defer r.mu.Unlock() + id, _ := r.nextDataIDLocked(defaultFileScope, nil) + return id +} + +// reserveDataID allocates a DataID before an open request is sent. The +// reservation prevents another concurrent open from selecting the same ID +// while the caller is still constructing and registering its bulk handle. +func (r *bulkRuntime) reserveDataID(scope string, requested uint64) (uint64, error) { + if r == nil { + return 0, errBulkRuntimeNil + } + scope = normalizeFileScope(scope) + r.mu.Lock() + defer r.mu.Unlock() + return r.nextDataIDLocked(scope, &requested) +} + +func (r *bulkRuntime) releaseDataID(scope string, dataID uint64) { + if r == nil || dataID == 0 { + return + } + scope = normalizeFileScope(scope) + r.mu.Lock() + defer r.mu.Unlock() + if reserved := r.reserved[scope]; reserved != nil { + delete(reserved, dataID) + if len(reserved) == 0 { + delete(r.reserved, scope) + } + } +} + +// nextDataIDLocked chooses an ID under r.mu. requested is nil for an +// internal auto allocation and points to zero/non-zero for a caller that +// wants a reservation for an outbound open. +func (r *bulkRuntime) nextDataIDLocked(scope string, requested *uint64) (uint64, error) { + if r == nil { + return 0, errBulkRuntimeNil + } + var wanted uint64 + if requested != nil { + wanted = *requested + } + dataScope := r.outbound[scope] + inboundScope := r.inbound[scope] + reserved := r.reserved[scope] + if wanted != 0 { + if !r.localDataID(wanted) { + return 0, errBulkDataIDEmpty + } + if dataScope != nil { + if _, exists := dataScope[wanted]; exists { + return 0, errBulkAlreadyExists + } + } + if inboundScope != nil { + if _, exists := inboundScope[wanted]; exists { + return 0, errBulkAlreadyExists + } + } + if _, exists := reserved[wanted]; exists { + return 0, errBulkAlreadyExists + } + if wanted > r.dataSeq { + r.dataSeq = wanted + } + if reserved == nil { + reserved = make(map[uint64]struct{}) + r.reserved[scope] = reserved + } + reserved[wanted] = struct{}{} + return wanted, nil + } + for { + candidate, ok := r.nextDataCandidateLocked() + if !ok { + return 0, errBulkDataIDExhausted + } + if dataScope != nil { + if _, exists := dataScope[candidate]; exists { + continue + } + } + if inboundScope != nil { + if _, exists := inboundScope[candidate]; exists { + continue + } + } + if _, exists := reserved[candidate]; exists { + continue + } + if requested != nil { + if reserved == nil { + reserved = make(map[uint64]struct{}) + r.reserved[scope] = reserved + } + reserved[candidate] = struct{}{} + } + return candidate, nil + } +} + +func (r *bulkRuntime) nextDataCandidateLocked() (uint64, bool) { + if r == nil { + return 0, false + } + step := r.dataStep + if step == 0 { + step = 1 + } + if r.dataSeq == 0 { + candidate := r.dataStart + if candidate == 0 { + return 0, false + } + r.dataSeq = candidate + return candidate, true + } + if r.dataSeq > ^uint64(0)-step { + return 0, false + } + candidate := r.dataSeq + step + if step == 2 && candidate%2 != r.dataStart%2 { + if candidate == ^uint64(0) { + return 0, false + } + candidate++ + } + r.dataSeq = candidate + return candidate, true +} + +func (r *bulkRuntime) localDataID(dataID uint64) bool { + if r == nil || dataID == 0 { + return false + } + if r.dataStep != 2 { + return true + } + return dataID%2 == r.dataStart%2 } func (r *bulkRuntime) setHandler(fn func(BulkAcceptInfo) error) { @@ -60,6 +227,46 @@ func (r *bulkRuntime) handlerSnapshot() func(BulkAcceptInfo) error { } func (r *bulkRuntime) register(scope string, bulk *bulkHandle) error { + // Keep the old helper usable by package-local callers and tests. Production + // paths use registerInbound/registerOutbound so a DataID can exist once in + // each direction without making inbound frame dispatch ambiguous. + return r.registerWithDirections(scope, bulk, bulkDataIndexBoth, false) +} + +func (r *bulkRuntime) registerInbound(scope string, bulk *bulkHandle) error { + return r.registerWithDirections(scope, bulk, bulkDataIndexInbound, false) +} + +// adoptInbound transfers ownership of a newly-created handle to the runtime. +// A failed registration is terminal because the handle has already started +// its background workers and must not be reused by the caller. +func (r *bulkRuntime) adoptInbound(scope string, bulk *bulkHandle) error { + err := r.registerInbound(scope, bulk) + if err != nil && bulk != nil { + bulk.markReset(err) + } + return err +} + +func (r *bulkRuntime) registerOutbound(scope string, bulk *bulkHandle) error { + return r.registerWithDirections(scope, bulk, bulkDataIndexOutbound, false) +} + +func (r *bulkRuntime) registerReserved(scope string, bulk *bulkHandle) error { + return r.registerWithDirections(scope, bulk, bulkDataIndexOutbound, true) +} + +// adoptReserved is the outbound counterpart of adoptInbound. The caller must +// still release an unconsumed DataID reservation when this method fails. +func (r *bulkRuntime) adoptReserved(scope string, bulk *bulkHandle) error { + err := r.registerReserved(scope, bulk) + if err != nil && bulk != nil { + bulk.markReset(err) + } + return err +} + +func (r *bulkRuntime) registerWithDirections(scope string, bulk *bulkHandle, direction bulkDataIndexDirection, consumeReservation bool) error { if r == nil { return errBulkRuntimeNil } @@ -73,19 +280,73 @@ func (r *bulkRuntime) register(scope string, bulk *bulkHandle) error { if _, ok := r.bulks[key]; ok { return errBulkAlreadyExists } - if bulk.dataID == 0 { + if direction == 0 { return errBulkDataIDEmpty } - dataScope := r.data[scope] - if dataScope == nil { - dataScope = make(map[uint64]*bulkHandle) - r.data[scope] = dataScope + if bulk.dataID == 0 { + dataID, err := r.nextDataIDLocked(scope, nil) + if err != nil { + return err + } + bulk.dataID = dataID + } else if direction&bulkDataIndexOutbound != 0 && r.localDataID(bulk.dataID) && bulk.dataID > r.dataSeq { + r.dataSeq = bulk.dataID } - if _, ok := dataScope[bulk.dataID]; ok { - return errBulkAlreadyExists + inbound := r.inbound[scope] + outbound := r.outbound[scope] + if direction&bulkDataIndexInbound != 0 && inbound != nil { + if _, ok := inbound[bulk.dataID]; ok { + return errBulkAlreadyExists + } + } + if direction&bulkDataIndexOutbound != 0 && outbound != nil { + if _, ok := outbound[bulk.dataID]; ok { + return errBulkAlreadyExists + } + } + // New peers use disjoint odd/even namespaces. Reject a legacy peer's + // colliding explicit ID when both directions would otherwise share one + // wire DataID; the inbound-frame router cannot disambiguate that case. + if r.dataStep == 2 { + if direction&bulkDataIndexInbound != 0 && outbound != nil { + if _, ok := outbound[bulk.dataID]; ok { + return errBulkAlreadyExists + } + } + if direction&bulkDataIndexOutbound != 0 && inbound != nil { + if _, ok := inbound[bulk.dataID]; ok { + return errBulkAlreadyExists + } + } + } + if direction&bulkDataIndexOutbound != 0 { + if reserved := r.reserved[scope]; reserved != nil { + if _, exists := reserved[bulk.dataID]; exists { + if !consumeReservation { + return errBulkAlreadyExists + } + delete(reserved, bulk.dataID) + if len(reserved) == 0 { + delete(r.reserved, scope) + } + } + } } r.bulks[key] = bulk - dataScope[bulk.dataID] = bulk + if direction&bulkDataIndexInbound != 0 { + if inbound == nil { + inbound = make(map[uint64]*bulkHandle) + r.inbound[scope] = inbound + } + inbound[bulk.dataID] = bulk + } + if direction&bulkDataIndexOutbound != 0 { + if outbound == nil { + outbound = make(map[uint64]*bulkHandle) + r.outbound[scope] = outbound + } + outbound[bulk.dataID] = bulk + } return nil } @@ -101,33 +362,141 @@ func (r *bulkRuntime) lookup(scope string, bulkID string) (*bulkHandle, bool) { } func (r *bulkRuntime) lookupByDataID(scope string, dataID uint64) (*bulkHandle, bool) { + return r.lookupByDataIDDirection(scope, dataID, bulkDataIndexBoth) +} + +func (r *bulkRuntime) lookupInboundByDataID(scope string, dataID uint64) (*bulkHandle, bool) { + return r.lookupByDataIDDirection(scope, dataID, bulkDataIndexInbound) +} + +// lookupInboundFrame chooses the local-open or peer-open index from the +// allocator partition. Both bulk kinds are duplex; the partition identifies +// which handle owns a wire DataID before consulting the corresponding map. +func (r *bulkRuntime) lookupInboundFrame(scope string, dataID uint64) (*bulkHandle, bool) { + if r == nil { + return nil, false + } + if r.dataStep != 2 { + return r.lookupByDataID(scope, dataID) + } + if r.localDataID(dataID) { + if bulk, ok := r.lookupOutboundByDataID(scope, dataID); ok { + return bulk, true + } + // Legacy peers may send DataID=0 in the open request. The receiver + // allocates an ID locally in that case, so accept the inbound index + // when the preferred outbound slot is absent. + return r.lookupInboundByDataID(scope, dataID) + } + if bulk, ok := r.lookupInboundByDataID(scope, dataID); ok { + return bulk, true + } + // Keep old peers that selected the local namespace routable when no + // inbound handle occupies the ID. + return r.lookupOutboundByDataID(scope, dataID) +} + +func (r *bulkRuntime) lookupOutboundByDataID(scope string, dataID uint64) (*bulkHandle, bool) { + return r.lookupByDataIDDirection(scope, dataID, bulkDataIndexOutbound) +} + +func (r *bulkRuntime) lookupByDataIDDirection(scope string, dataID uint64, direction bulkDataIndexDirection) (*bulkHandle, bool) { if r == nil || dataID == 0 { return nil, false } scope = normalizeFileScope(scope) r.mu.RLock() defer r.mu.RUnlock() - dataScope := r.data[scope] - if dataScope == nil { + var inbound, outbound *bulkHandle + if direction&bulkDataIndexInbound != 0 { + if dataScope := r.inbound[scope]; dataScope != nil { + inbound = dataScope[dataID] + } + } + if direction&bulkDataIndexOutbound != 0 { + if dataScope := r.outbound[scope]; dataScope != nil { + outbound = dataScope[dataID] + } + } + if direction == bulkDataIndexBoth && inbound != nil && outbound != nil && inbound != outbound { return nil, false } - bulk, ok := dataScope[dataID] - return bulk, ok + if inbound != nil { + return inbound, true + } + if outbound != nil { + return outbound, true + } + return nil, false } -func (r *bulkRuntime) remove(scope string, bulkID string) { - if r == nil || bulkID == "" { +// lookupControl resolves the identity carried by a control message. A +// supplied BulkID is authoritative and, when present, DataID must agree. A +// DataID-only message is accepted only when it maps to one direction; using a +// colliding ID from the other direction would otherwise reset the wrong bulk. +func (r *bulkRuntime) lookupControl(scope string, bulkID string, dataID uint64) (*bulkHandle, bool) { + if r == nil { + return nil, false + } + scope = normalizeFileScope(scope) + r.mu.RLock() + defer r.mu.RUnlock() + if bulkID != "" { + bulk, ok := r.bulks[bulkRuntimeKey(scope, bulkID)] + if !ok || bulk == nil { + return nil, false + } + if dataID != 0 && bulk.dataID != dataID { + return nil, false + } + return bulk, true + } + if dataID == 0 { + return nil, false + } + var inbound, outbound *bulkHandle + if dataScope := r.inbound[scope]; dataScope != nil { + inbound = dataScope[dataID] + } + if dataScope := r.outbound[scope]; dataScope != nil { + outbound = dataScope[dataID] + } + if inbound != nil && outbound != nil && inbound != outbound { + return nil, false + } + if inbound != nil { + return inbound, true + } + return outbound, outbound != nil +} + +func (r *bulkRuntime) remove(scope string, expected *bulkHandle) { + if r == nil || expected == nil || expected.id == "" { return } scope = normalizeFileScope(scope) - key := bulkRuntimeKey(scope, bulkID) + key := bulkRuntimeKey(scope, expected.id) r.mu.Lock() defer r.mu.Unlock() - if bulk := r.bulks[key]; bulk != nil && bulk.dataID != 0 { - if dataScope := r.data[scope]; dataScope != nil { - delete(dataScope, bulk.dataID) + bulk := r.bulks[key] + if bulk != expected { + return + } + if bulk.dataID != 0 { + if dataScope := r.inbound[scope]; dataScope != nil { + if dataScope[bulk.dataID] == bulk { + delete(dataScope, bulk.dataID) + } if len(dataScope) == 0 { - delete(r.data, scope) + delete(r.inbound, scope) + } + } + if dataScope := r.outbound[scope]; dataScope != nil { + if dataScope[bulk.dataID] == bulk { + delete(dataScope, bulk.dataID) + } + if len(dataScope) == 0 { + delete(r.outbound, scope) } } } @@ -145,6 +514,44 @@ func (r *bulkRuntime) closeScope(scope string, err error) { }, err) } +func (r *bulkRuntime) closeClientRoute(route clientSessionRoute, err error) { + if r == nil { + return + } + if !r.mu.TryRLock() { + go r.closeClientRouteBlocking(route, err) + return + } + bulks := r.collectClientRouteLocked(route) + r.mu.RUnlock() + r.resetClientRouteHandles(bulks, err) +} + +func (r *bulkRuntime) closeClientRouteBlocking(route clientSessionRoute, err error) { + r.mu.RLock() + bulks := r.collectClientRouteLocked(route) + r.mu.RUnlock() + r.resetClientRouteHandles(bulks, err) +} + +func (r *bulkRuntime) collectClientRouteLocked(route clientSessionRoute) []*bulkHandle { + bulks := make([]*bulkHandle, 0) + for _, bulk := range r.bulks { + if bulk == nil || !sameClientSessionRoute(bulk.clientRoute, route) { + continue + } + bulks = append(bulks, bulk) + } + return bulks +} + +func (r *bulkRuntime) resetClientRouteHandles(bulks []*bulkHandle, err error) { + resetErr := bulkRuntimeCloseError(err) + for _, bulk := range bulks { + bulk.markReset(resetErr) + } +} + func (r *bulkRuntime) closeMatching(match func(string) bool, err error) { if r == nil || match == nil { return diff --git a/bulk_shared_batch.go b/bulk_shared_batch.go index 3f8e484..4510a87 100644 --- a/bulk_shared_batch.go +++ b/bulk_shared_batch.go @@ -49,6 +49,24 @@ func bulkFastBatchPlainLen(frames []bulkFastFrame) int { return total } +func bulkFastBatchPlainLenChecked(frames []bulkFastFrame) (int, error) { + if len(frames) == 0 || len(frames) > bulkFastBatchMaxItems { + return 0, errBulkFastPayloadInvalid + } + total := bulkFastBatchHeaderLen + for _, frame := range frames { + itemLen := bulkFastBatchFrameLen(frame) + if itemLen < bulkFastBatchItemHeaderLen || itemLen > bulkFastBatchMaxPlainBytes { + return 0, errBulkFastPayloadInvalid + } + if total > bulkFastBatchMaxPlainBytes-itemLen { + return 0, errBulkFastPayloadInvalid + } + total += itemLen + } + return total, nil +} + func encodeBulkFastFramePayload(frame bulkFastFrame) ([]byte, error) { return encodeBulkFastControlFrame(frame.Type, frame.Flags, frame.DataID, frame.Seq, frame.Payload) } @@ -82,10 +100,11 @@ func encodeBulkFastFramePayloadPooled(runtime *modernPSKCodecRuntime, frame bulk } func encodeBulkFastBatchPlain(frames []bulkFastFrame) ([]byte, error) { - if len(frames) == 0 { - return nil, errBulkFastPayloadInvalid + plainLen, err := bulkFastBatchPlainLenChecked(frames) + if err != nil { + return nil, err } - buf := make([]byte, bulkFastBatchPlainLen(frames)) + buf := make([]byte, plainLen) if err := writeBulkFastBatchPlain(buf, frames); err != nil { return nil, err } @@ -96,7 +115,10 @@ func encodeBulkFastBatchPayloadFast(encode transportFastPlainEncoder, secretKey if encode == nil { return nil, errTransportPayloadEncryptFailed } - plainLen := bulkFastBatchPlainLen(frames) + plainLen, err := bulkFastBatchPlainLenChecked(frames) + if err != nil { + return nil, err + } return encode(secretKey, plainLen, func(dst []byte) error { return writeBulkFastBatchPlain(dst, frames) }) @@ -106,13 +128,21 @@ func encodeBulkFastBatchPayloadPooled(runtime *modernPSKCodecRuntime, frames []b if runtime == nil { return nil, nil, errTransportPayloadEncryptFailed } - return runtime.sealFilledPayloadPooled(bulkFastBatchPlainLen(frames), func(dst []byte) error { + plainLen, err := bulkFastBatchPlainLenChecked(frames) + if err != nil { + return nil, nil, err + } + return runtime.sealFilledPayloadPooled(plainLen, func(dst []byte) error { return writeBulkFastBatchPlain(dst, frames) }) } func writeBulkFastBatchPlain(dst []byte, frames []bulkFastFrame) error { - if len(frames) == 0 || len(dst) != bulkFastBatchPlainLen(frames) { + plainLen, err := bulkFastBatchPlainLenChecked(frames) + if err != nil { + return err + } + if len(dst) != plainLen { return errBulkFastPayloadInvalid } copy(dst[:4], bulkFastBatchMagic) @@ -145,10 +175,11 @@ func walkBulkFastBatchPlain(payload []byte, fn func(bulkFastFrame) error) (bool, if payload[4] != bulkFastBatchVersion { return true, errBulkFastPayloadInvalid } - count := int(binary.BigEndian.Uint32(payload[8:12])) - if count <= 0 { + wireCount := binary.BigEndian.Uint32(payload[8:12]) + if wireCount == 0 || wireCount > bulkFastBatchMaxItems { return true, errBulkFastPayloadInvalid } + count := int(wireCount) offset := bulkFastBatchHeaderLen for index := 0; index < count; index++ { if len(payload)-offset < bulkFastBatchItemHeaderLen { @@ -163,11 +194,12 @@ func walkBulkFastBatchPlain(payload []byte, fn func(bulkFastFrame) error) (bool, flags := payload[offset+1] dataID := binary.BigEndian.Uint64(payload[offset+4 : offset+12]) seq := binary.BigEndian.Uint64(payload[offset+12 : offset+20]) - payloadLen := int(binary.BigEndian.Uint32(payload[offset+20 : offset+24])) + wirePayloadLen := binary.BigEndian.Uint32(payload[offset+20 : offset+24]) offset += bulkFastBatchItemHeaderLen - if dataID == 0 || payloadLen < 0 || len(payload)-offset < payloadLen { + if dataID == 0 || uint64(wirePayloadLen) > uint64(len(payload)-offset) { return true, errBulkFastPayloadInvalid } + payloadLen := int(wirePayloadLen) if fn != nil { if err := fn(bulkFastFrame{ Type: frameType, diff --git a/bulk_shared_batch_test.go b/bulk_shared_batch_test.go index f4dcf6e..cdeb2dc 100644 --- a/bulk_shared_batch_test.go +++ b/bulk_shared_batch_test.go @@ -2,10 +2,42 @@ package notify import ( "context" + "errors" + "sync/atomic" "testing" "time" ) +func TestBulkBatchSenderEncodeRequestsReleasesPayloadsOnError(t *testing.T) { + if !bulkFastPathSupportsSharedBatch(bulkFastPathVersionV2) { + t.Fatal("v2 should support shared batch") + } + var released atomic.Int32 + sender := &bulkBatchSender{ + codec: bulkBatchCodec{ + encodeSingle: func(frame bulkFastFrame) ([]byte, func(), error) { + return nil, nil, errors.New("single encode failed") + }, + encodeBatch: func(frames []bulkFastFrame) ([]byte, func(), error) { + return []byte("batch"), func() { released.Add(1) }, nil + }, + }, + } + _, err := sender.encodeRequests([]bulkBatchRequest{ + {frames: []bulkFastFrame{ + {Type: bulkFastPayloadTypeData, DataID: 1, Payload: []byte("a")}, + {Type: bulkFastPayloadTypeData, DataID: 1, Payload: []byte("a2")}, + }, fastPathVersion: bulkFastPathVersionV2}, + {frames: []bulkFastFrame{{Type: bulkFastPayloadTypeData, DataID: 1, Payload: []byte("b")}}, fastPathVersion: 1}, + }) + if err == nil { + t.Fatal("encodeRequests unexpectedly succeeded") + } + if got := released.Load(); got != 1 { + t.Fatalf("released payloads = %d, want 1", got) + } +} + func TestBulkFastBatchPlainRoundTrip(t *testing.T) { releasePayload, err := encodeBulkDedicatedReleasePayload(4096, 2) if err != nil { diff --git a/bulk_test.go b/bulk_test.go index deee2c7..b8d882e 100644 --- a/bulk_test.go +++ b/bulk_test.go @@ -11,6 +11,64 @@ import ( "time" ) +func TestFinalizedBulkRejectsDedicatedAttachOperations(t *testing.T) { + operations := []struct { + name string + attach func(*bulkHandle, net.Conn) error + }{ + { + name: "attach-owned", + attach: func(bulk *bulkHandle, conn net.Conn) error { + return bulk.attachDedicatedConn(conn) + }, + }, + { + name: "attach-shared", + attach: func(bulk *bulkHandle, conn net.Conn) error { + return bulk.attachDedicatedConnShared(conn) + }, + }, + { + name: "replace-owned", + attach: func(bulk *bulkHandle, conn net.Conn) error { + _, _, err := bulk.replaceDedicatedConn(conn) + return err + }, + }, + { + name: "replace-shared", + attach: func(bulk *bulkHandle, conn net.Conn) error { + _, _, err := bulk.replaceDedicatedConnShared(conn) + return err + }, + }, + } + + for _, operation := range operations { + t.Run(operation.name, func(t *testing.T) { + bulk := newBulkHandle(context.Background(), nil, clientFileScope(), BulkOpenRequest{ + BulkID: "finalized-dedicated-attach", + DataID: 1, + Dedicated: true, + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.finalize() + + left, right := net.Pipe() + defer left.Close() + defer right.Close() + if err := operation.attach(bulk, left); !errors.Is(err, io.ErrClosedPipe) { + t.Fatalf("attach after finalize error = %v, want %v", err, io.ErrClosedPipe) + } + if got := bulk.dedicatedConnSnapshot(); got != nil { + t.Fatalf("finalized bulk retained dedicated conn %v", got) + } + if got := bulk.dedicatedAttachStateSnapshot(); got != bulkDedicatedAttachStateClosed { + t.Fatalf("dedicated state = %v, want closed", got) + } + }) + } +} + func TestBulkOpenRoundTripTCP(t *testing.T) { server := NewServer().(*ServerCommon) if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { @@ -293,6 +351,87 @@ func TestDedicatedBulkOpenUnblocksOnBlockingFirstWrite(t *testing.T) { waitForBulkContextDone(t, bulk.Context(), 2*time.Second) } +func TestSharedBulkOpenRoutesHandlerWriteBeforeOpenReply(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKServer failed: %v", err) + } + payload := "shared-server-first-write" + server.SetBulkHandler(func(info BulkAcceptInfo) error { + if _, err := io.WriteString(info.Bulk, payload); err != nil { + return err + } + return nil + }) + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatalf("server Listen failed: %v", err) + } + defer func() { _ = server.Stop() }() + + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKClient failed: %v", err) + } + if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { + t.Fatalf("client Connect failed: %v", err) + } + defer func() { _ = client.Stop() }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + bulk, err := client.OpenBulk(ctx, BulkOpenOptions{ + ID: "shared-handler-first-write", + Range: BulkRange{Length: int64(len(payload))}, + Mode: BulkOpenModeShared, + }) + if err != nil { + t.Fatalf("client OpenBulk failed: %v", err) + } + defer bulk.Close() + readBulkExactly(t, bulk, payload, 2*time.Second) +} + +func TestServerSharedBulkOpenRoutesHandlerWriteBeforeOpenReply(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKServer failed: %v", err) + } + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatalf("server Listen failed: %v", err) + } + defer func() { _ = server.Stop() }() + + payload := "shared-client-first-write" + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKClient failed: %v", err) + } + client.SetBulkHandler(func(info BulkAcceptInfo) error { + if _, err := io.WriteString(info.Bulk, payload); err != nil { + return err + } + return nil + }) + if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { + t.Fatalf("client Connect failed: %v", err) + } + defer func() { _ = client.Stop() }() + + logical := waitForTransferControlLogicalConn(t, server, 2*time.Second) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + bulk, err := server.OpenBulkLogical(ctx, logical, BulkOpenOptions{ + ID: "server-shared-handler-first-write", + Range: BulkRange{Length: int64(len(payload))}, + Mode: BulkOpenModeShared, + }) + if err != nil { + t.Fatalf("server OpenBulkLogical failed: %v", err) + } + defer bulk.Close() + readBulkExactly(t, bulk, payload, 2*time.Second) +} + func TestServerOpenBulkLogicalDedicatedUnblocksOnBlockingFirstRead(t *testing.T) { server := NewServer().(*ServerCommon) if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { diff --git a/client.go b/client.go index d437c67..96319fe 100644 --- a/client.go +++ b/client.go @@ -70,6 +70,8 @@ type ClientCommon struct { streamRuntime *streamRuntime recordRuntime *recordRuntime bulkRuntime *bulkRuntime + bulkRecovery *bulkRecoveryQueue + bulkRecoveryMu sync.Mutex bulkDefaultOpenMode BulkOpenMode bulkNetworkProfile BulkNetworkProfile bulkOpenTuning BulkOpenTuning @@ -139,6 +141,7 @@ func NewClient() Client { client.streamRuntime = newStreamRuntime("cstrm") client.recordRuntime = newRecordRuntime() client.bulkRuntime = newBulkRuntime("cblk") + client.bulkRecovery = newBulkRecoveryQueue(client.reportBulkRecoveryError) client.bulkDedicatedLanes = make(map[uint32]*bulkDedicatedLane) if client.bulkDedicatedAttachLimit > 0 { client.bulkDedicatedAttachSem = make(chan struct{}, client.bulkDedicatedAttachLimit) diff --git a/client_bulk.go b/client_bulk.go index e4d636b..dc0a83b 100644 --- a/client_bulk.go +++ b/client_bulk.go @@ -3,8 +3,11 @@ package notify import ( "context" "errors" + "time" ) +const bulkOpenRecoveryTimeout = 2 * time.Second + func (c *ClientCommon) SetBulkHandler(fn func(BulkAcceptInfo) error) { runtime := c.getBulkRuntime() if runtime == nil { @@ -33,21 +36,21 @@ func (c *ClientCommon) OpenBulk(ctx context.Context, opt BulkOpenOptions) (Bulk, switch opt.Mode { case BulkOpenModeDedicated: opt.Dedicated = true - return c.openBulkWithDedicatedMode(ctx, opt) + return c.openBulkWithDedicatedMode(ctx, opt, false) case BulkOpenModeAuto: // Auto mode prefers dedicated path and falls back to shared if dedicated fails. if err := clientDedicatedBulkSupportError(c); err == nil { dedicatedOpt := opt dedicatedOpt.Mode = BulkOpenModeDedicated dedicatedOpt.Dedicated = true - bulk, dedicatedErr := c.openBulkWithDedicatedMode(ctx, dedicatedOpt) + bulk, dedicatedErr := c.openBulkWithDedicatedMode(ctx, dedicatedOpt, true) if dedicatedErr == nil { return bulk, nil } sharedOpt := opt sharedOpt.Mode = BulkOpenModeShared sharedOpt.Dedicated = false - sharedBulk, sharedErr := c.openBulkWithDedicatedMode(ctx, sharedOpt) + sharedBulk, sharedErr := c.openBulkWithDedicatedMode(ctx, sharedOpt, false) if sharedErr == nil { c.bulkAttachFallbackCount.Add(1) return sharedBulk, nil @@ -57,22 +60,26 @@ func (c *ClientCommon) OpenBulk(ctx context.Context, opt BulkOpenOptions) (Bulk, opt.Mode = BulkOpenModeShared opt.Dedicated = false c.bulkAttachFallbackCount.Add(1) - return c.openBulkWithDedicatedMode(ctx, opt) + return c.openBulkWithDedicatedMode(ctx, opt, false) case BulkOpenModeShared, BulkOpenModeDefault: opt.Mode = BulkOpenModeShared opt.Dedicated = false - return c.openBulkWithDedicatedMode(ctx, opt) + return c.openBulkWithDedicatedMode(ctx, opt, false) default: opt.Mode = BulkOpenModeShared opt.Dedicated = false - return c.openBulkWithDedicatedMode(ctx, opt) + return c.openBulkWithDedicatedMode(ctx, opt, false) } } -func (c *ClientCommon) openBulkWithDedicatedMode(ctx context.Context, opt BulkOpenOptions) (Bulk, error) { +func (c *ClientCommon) openBulkWithDedicatedMode(ctx context.Context, opt BulkOpenOptions, waitForReset bool) (Bulk, error) { if c == nil { return nil, errBulkClientNil } + route := c.clientSessionRouteSnapshot() + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return nil, err + } opt = applyBulkOpenTuningDefaults(opt, c.bulkOpenTuningSnapshot()) runtime := c.getBulkRuntime() if runtime == nil { @@ -90,112 +97,199 @@ func (c *ClientCommon) openBulkWithDedicatedMode(ctx context.Context, opt BulkOp if !validBulkRange(req.Range) { return nil, errBulkRangeInvalid } - if _, exists := runtime.lookup(clientFileScope(), req.BulkID); exists { - return nil, errBulkAlreadyExists + if existing, exists := runtime.lookup(clientFileScope(), req.BulkID); exists { + if existing.acceptsClientSessionRoute(route) { + return nil, errBulkAlreadyExists + } + existing.markReset(errTransportDetached) + } + if req.DataID == 0 { + var reserveErr error + req.DataID, reserveErr = runtime.reserveDataID(clientFileScope(), 0) + if reserveErr != nil { + return nil, reserveErr + } } if req.Dedicated { - req.DedicatedLaneID = c.reserveBulkDedicatedLane() - if req.DataID == 0 { - req.DataID = runtime.nextDataID() + var laneErr error + req.DedicatedLaneID, laneErr = c.reserveBulkDedicatedLaneAtRoute(route) + if laneErr != nil { + runtime.releaseDataID(clientFileScope(), req.DataID) + return nil, laneErr } if req.AttachToken == "" { req.AttachToken = newBulkAttachToken() } - bulk := newBulkHandle(c.clientStopContextSnapshot(), runtime, clientFileScope(), req, c.currentClientSessionEpoch(), nil, nil, 0, clientBulkCloseSender(c), clientBulkResetSender(c), clientBulkDataSender(c, c.currentClientSessionEpoch()), clientBulkWriteSender(c, c.currentClientSessionEpoch()), clientBulkReleaseSender(c)) + bulk := newBulkHandle(clientSessionRouteContext(route), runtime, clientFileScope(), req, route.epoch, nil, nil, 0, clientBulkCloseSender(c), clientBulkResetSender(c), clientBulkDataSender(c, route), clientBulkWriteSender(c, route), clientBulkReleaseSender(c)) bulk.setClientSnapshotOwner(c) + bulk.setClientSessionRoute(route) bulk.markAcceptHandled() - if err := runtime.register(clientFileScope(), bulk); err != nil { - c.releaseBulkDedicatedLane(req.DedicatedLaneID) + bulk.markDedicatedLaneReserved() + if err := runtime.adoptReserved(clientFileScope(), bulk); err != nil { + runtime.releaseDataID(clientFileScope(), req.DataID) return nil, err } - resp, err := sendBulkOpenClient(ctx, c, req) + resp, err := sendBulkOpenClientAtRoute(ctx, c, route, req) if err != nil { + runtime.releaseDataID(clientFileScope(), req.DataID) + cleanupErr := c.cleanupBulkResetAtRoute(ctx, route, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: err.Error()}, waitForReset) bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } if resp.DataID != 0 && resp.DataID != req.DataID { err = errBulkAlreadyExists - _, _ = sendBulkResetClient(context.Background(), c, BulkResetRequest{ + cleanupErr := c.cleanupBulkResetAtRoute(ctx, route, BulkResetRequest{ BulkID: req.BulkID, - DataID: req.DataID, Error: "bulk dedicated data id mismatch", - }) + }, waitForReset) bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } if resp.TransportGeneration != 0 { - bulk.transportGeneration = resp.TransportGeneration + bulk.setTransportGeneration(resp.TransportGeneration) } if resp.FastPathVersion != 0 { - bulk.fastPathVersion = normalizeBulkFastPathVersion(resp.FastPathVersion) + bulk.setFastPathVersion(resp.FastPathVersion) } if resp.AttachToken != "" { req.AttachToken = resp.AttachToken bulk.setDedicatedAttachToken(resp.AttachToken) } if err := c.attachDedicatedBulkSidecar(ctx, bulk); err != nil { - _, _ = sendBulkResetClient(context.Background(), c, BulkResetRequest{ + cleanupErr := c.cleanupBulkResetAtRoute(ctx, route, BulkResetRequest{ BulkID: req.BulkID, DataID: req.DataID, Error: err.Error(), - }) + }, waitForReset) bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } if err := bulk.waitAcceptReady(ctx); err != nil { - _, _ = sendBulkResetClient(context.Background(), c, BulkResetRequest{ - BulkID: req.BulkID, - DataID: req.DataID, - Error: err.Error(), - }) + var cleanupErr error + if bulk.resetErrSnapshot() == nil { + cleanupErr = c.cleanupBulkResetAtRoute(ctx, route, BulkResetRequest{ + BulkID: req.BulkID, + DataID: req.DataID, + Error: err.Error(), + }, waitForReset) + } else { + // A ready error already reset the remote handle. Keep the old + // asynchronous cleanup for compatibility with that path and avoid + // racing a concurrent dedicated attach teardown. + c.bestEffortBulkResetAtRoute(route, BulkResetRequest{ + BulkID: req.BulkID, + DataID: req.DataID, + Error: err.Error(), + }) + } bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } return bulk, nil } - resp, err := sendBulkOpenClient(ctx, c, req) - if err != nil { + bulk := newBulkHandle(clientSessionRouteContext(route), runtime, clientFileScope(), req, route.epoch, nil, nil, 0, clientBulkCloseSender(c), clientBulkResetSender(c), clientBulkDataSender(c, route), clientBulkWriteSender(c, route), clientBulkReleaseSender(c)) + bulk.setClientSnapshotOwner(c) + bulk.setClientSessionRoute(route) + bulk.markAcceptHandled() + if err := runtime.adoptReserved(clientFileScope(), bulk); err != nil { + runtime.releaseDataID(clientFileScope(), req.DataID) return nil, err } - if resp.DataID != 0 { - req.DataID = resp.DataID + resp, err := sendBulkOpenClientAtRoute(ctx, c, route, req) + if err != nil { + c.bestEffortBulkResetAtRoute(route, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: err.Error()}) + bulk.markReset(err) + return nil, err + } + if resp.DataID != 0 && resp.DataID != req.DataID { + err = errBulkAlreadyExists + c.bestEffortBulkResetAtRoute(route, BulkResetRequest{BulkID: req.BulkID, Error: "bulk data id mismatch"}) + bulk.markReset(err) + return nil, err } if resp.FastPathVersion != 0 { - req.FastPathVersion = resp.FastPathVersion + bulk.setFastPathVersion(resp.FastPathVersion) } - req.Dedicated = resp.Dedicated - if resp.AttachToken != "" { - req.AttachToken = resp.AttachToken - } - if req.DataID == 0 { - return nil, errBulkDataIDEmpty - } - bulk := newBulkHandle(c.clientStopContextSnapshot(), runtime, clientFileScope(), req, c.currentClientSessionEpoch(), nil, nil, resp.TransportGeneration, clientBulkCloseSender(c), clientBulkResetSender(c), clientBulkDataSender(c, c.currentClientSessionEpoch()), clientBulkWriteSender(c, c.currentClientSessionEpoch()), clientBulkReleaseSender(c)) - bulk.setClientSnapshotOwner(c) - bulk.markAcceptHandled() - if err := runtime.register(clientFileScope(), bulk); err != nil { - c.releaseBulkDedicatedLane(req.DedicatedLaneID) - _, _ = sendBulkResetClient(context.Background(), c, BulkResetRequest{ - BulkID: req.BulkID, - DataID: req.DataID, - Error: err.Error(), - }) + if resp.Dedicated { + err = errBulkRejected + c.bestEffortBulkResetAtRoute(route, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: "shared bulk upgraded to dedicated"}) + bulk.markReset(err) return nil, err } - if bulk.Dedicated() { - if err := c.attachDedicatedBulkSidecar(ctx, bulk); err != nil { - runtime.remove(clientFileScope(), bulk.ID()) - _, _ = sendBulkResetClient(context.Background(), c, BulkResetRequest{ - BulkID: bulk.ID(), - DataID: bulk.dataIDSnapshot(), - Error: err.Error(), - }) - return nil, err - } + if resp.AttachToken != "" { + bulk.setDedicatedAttachToken(resp.AttachToken) } + bulk.setTransportGeneration(resp.TransportGeneration) return bulk, nil } +func (c *ClientCommon) bestEffortBulkReset(req BulkResetRequest) { + if c == nil { + return + } + c.bestEffortBulkResetAtRoute(c.clientSessionRouteSnapshot(), req) +} + +func (c *ClientCommon) bestEffortBulkResetAtEpoch(epoch uint64, req BulkResetRequest) { + route := c.clientSessionRouteSnapshot() + route.epoch = epoch + c.bestEffortBulkResetAtRoute(route, req) +} + +func (c *ClientCommon) bestEffortBulkResetAtRoute(route clientSessionRoute, req BulkResetRequest) { + if c == nil { + return + } + task := newClientBulkResetRecoveryTaskAtRoute(c, route, req) + q := c.bulkRecoveryQueue() + if !q.enqueue(task) { + c.handleBulkRecoveryOverflowAtRoute(route, req) + } +} + +func (c *ClientCommon) bulkRecoveryQueue() *bulkRecoveryQueue { + if c == nil { + return nil + } + c.bulkRecoveryMu.Lock() + defer c.bulkRecoveryMu.Unlock() + if c.bulkRecovery == nil { + c.bulkRecovery = newBulkRecoveryQueue(c.reportBulkRecoveryError) + } + return c.bulkRecovery +} + +func newClientBulkResetRecoveryTask(c *ClientCommon, epoch uint64, req BulkResetRequest) bulkRecoveryTask { + route := c.clientSessionRouteSnapshot() + route.epoch = epoch + return newClientBulkResetRecoveryTaskAtRoute(c, route, req) +} + +func newClientBulkResetRecoveryTaskAtRoute(c *ClientCommon, route clientSessionRoute, req BulkResetRequest) bulkRecoveryTask { + return func(ctx context.Context) error { + if !c.clientSessionRouteCurrent(route) { + return transportDetachedSessionEpochError() + } + _, err := sendBulkResetClientAtRoute(ctx, c, route, req) + if errors.Is(err, errBulkNotFound) { + return nil + } + return err + } +} + func clientBulkRequest(runtime *bulkRuntime, opt BulkOpenOptions) BulkOpenRequest { opt = normalizeBulkOpenOptions(opt) id := opt.ID @@ -224,8 +318,9 @@ func clientBulkCloseSender(c *ClientCommon) bulkCloseSender { } return c.sendDedicatedBulkClose(ctx, bulk, full) } - _, err := sendBulkCloseClient(ctx, c, BulkCloseRequest{ + _, err := sendBulkCloseClientAtRoute(ctx, c, bulk.clientSessionRouteSnapshot(), BulkCloseRequest{ BulkID: bulk.ID(), + DataID: bulk.dataIDSnapshot(), Full: full, }) return err @@ -240,7 +335,7 @@ func clientBulkResetSender(c *ClientCommon) bulkResetSender { } return c.sendDedicatedBulkReset(ctx, bulk, message) } - _, err := sendBulkResetClient(ctx, c, BulkResetRequest{ + _, err := sendBulkResetClientAtRoute(ctx, c, bulk.clientSessionRouteSnapshot(), BulkResetRequest{ BulkID: bulk.ID(), DataID: bulk.dataIDSnapshot(), Error: message, @@ -249,7 +344,7 @@ func clientBulkResetSender(c *ClientCommon) bulkResetSender { } } -func clientBulkDataSender(c *ClientCommon, epoch uint64) bulkDataSender { +func clientBulkDataSender(c *ClientCommon, route clientSessionRoute) bulkDataSender { return func(ctx context.Context, bulk *bulkHandle, chunk []byte) error { if c == nil { return errBulkClientNil @@ -267,18 +362,18 @@ func clientBulkDataSender(c *ClientCommon, epoch uint64) bulkDataSender { } return c.sendDedicatedBulkData(ctx, bulk, chunk) } - if epoch != 0 && !c.isClientSessionEpochCurrent(epoch) { + if !c.clientSessionRouteCurrent(route) { return errTransportDetached } dataID := bulk.dataIDSnapshot() if dataID == 0 { return errBulkDataPathNotReady } - return c.sendFastBulkData(ctx, dataID, bulk.nextOutboundDataSeq(), chunk, bulk.fastPathVersionSnapshot()) + return c.sendFastBulkDataAtRoute(ctx, route, dataID, bulk.nextOutboundDataSeq(), chunk, bulk.fastPathVersionSnapshot()) } } -func clientBulkWriteSender(c *ClientCommon, epoch uint64) bulkWriteSender { +func clientBulkWriteSender(c *ClientCommon, route clientSessionRoute) bulkWriteSender { return func(ctx context.Context, bulk *bulkHandle, startSeq uint64, payload []byte, payloadOwned bool) (int, error) { if c == nil { return 0, errBulkClientNil @@ -296,7 +391,7 @@ func clientBulkWriteSender(c *ClientCommon, epoch uint64) bulkWriteSender { } return c.sendDedicatedBulkWrite(ctx, bulk, startSeq, payload, payloadOwned) } - if epoch != 0 && !c.isClientSessionEpochCurrent(epoch) { + if !c.clientSessionRouteCurrent(route) { return 0, errTransportDetached } if bulk == nil { @@ -306,11 +401,12 @@ func clientBulkWriteSender(c *ClientCommon, epoch uint64) bulkWriteSender { if dataID == 0 { return 0, errBulkDataPathNotReady } - return c.sendFastBulkWrite(ctx, dataID, startSeq, bulk.chunkSize, bulk.fastPathVersionSnapshot(), payload, payloadOwned) + return c.sendFastBulkWriteAtRoute(ctx, route, dataID, startSeq, bulk.chunkSize, bulk.fastPathVersionSnapshot(), payload, payloadOwned) } } func clientBulkReleaseSender(c *ClientCommon) bulkReleaseSender { + fallbackRoute := c.clientSessionRouteSnapshot() return func(bulk *bulkHandle, bytes int64, chunks int) error { if c == nil || bulk == nil { return errBulkClientNil @@ -326,14 +422,18 @@ func clientBulkReleaseSender(c *ClientCommon) bulkReleaseSender { if bulk.Dedicated() { return c.sendDedicatedBulkRelease(ctx, bulk, bytes, chunks) } + route := bulk.clientSessionRouteSnapshot() + if !route.bound() { + route = fallbackRoute + } if bulk.fastPathVersionSnapshot() >= bulkFastPathVersionV2 { payload, err := encodeBulkDedicatedReleasePayload(bytes, chunks) if err != nil { return err } - return c.sendFastBulkControl(ctx, bulkFastPayloadTypeRelease, 0, bulk.dataIDSnapshot(), 0, bulk.fastPathVersionSnapshot(), payload) + return c.sendFastBulkControlAtRoute(ctx, route, bulkFastPayloadTypeRelease, 0, bulk.dataIDSnapshot(), 0, bulk.fastPathVersionSnapshot(), payload) } - return sendBulkReleaseClient(ctx, c, BulkReleaseRequest{ + return sendBulkReleaseClientAtRoute(ctx, c, route, BulkReleaseRequest{ BulkID: bulk.ID(), DataID: bulk.dataIDSnapshot(), Bytes: bytes, diff --git a/client_config.go b/client_config.go index 64c5645..999fca4 100644 --- a/client_config.go +++ b/client_config.go @@ -12,6 +12,8 @@ func (c *ClientCommon) DebugMode(dmg bool) { } func (c *ClientCommon) IsDebugMode() bool { + c.mu.Lock() + defer c.mu.Unlock() return c.debugMode } diff --git a/client_conn.go b/client_conn.go index 1939f4b..c6ce427 100644 --- a/client_conn.go +++ b/client_conn.go @@ -2,6 +2,7 @@ package notify import ( "b612.me/starcrypto" + "b612.me/stario" "fmt" "net" "sync/atomic" @@ -73,7 +74,7 @@ func (c *ClientConn) readTUMessageLoop(rt *clientConnSessionRuntime) { generation := rt.transportGeneration defer closeClientConnSessionRuntimeTransportDone(rt) if conn != nil && !isPacketTransportConn(conn) { - reader := newTransportFrameReader(conn, nil) + reader := newTransportFrameReader(conn, stario.NewQueueCtx(stopCtx, 4, transportFrameMaxPayloadBytes)) for { select { case <-sessionStopChan(stopCtx): diff --git a/client_conn_transport.go b/client_conn_transport.go index c94f724..2e4b053 100644 --- a/client_conn_transport.go +++ b/client_conn_transport.go @@ -43,7 +43,7 @@ func (c *LogicalConn) readTUMessageLoop(rt *clientConnSessionRuntime) { generation := rt.transportGeneration defer closeClientConnSessionRuntimeTransportDone(rt) if conn != nil && !isPacketTransportConn(conn) { - reader := newTransportFrameReader(conn, nil) + reader := newTransportFrameReader(conn, stario.NewQueueCtx(stopCtx, 4, transportFrameMaxPayloadBytes)) for { select { case <-sessionStopChan(stopCtx): diff --git a/client_runtime.go b/client_runtime.go index 28500fe..a0ca92f 100644 --- a/client_runtime.go +++ b/client_runtime.go @@ -5,7 +5,6 @@ import ( "context" "errors" "fmt" - "math" "net" "sync/atomic" "time" @@ -286,7 +285,7 @@ func (c *ClientCommon) startClientWithConn(conn net.Conn) error { func (c *ClientCommon) startClientWithConnSource(conn net.Conn, source *clientConnectSource) error { stopCtx, stopFn := context.WithCancel(context.Background()) epoch := c.beginClientSessionEpoch() - queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) + queue := stario.NewQueueCtx(stopCtx, 4, transportFrameMaxPayloadBytes) c.setClientConnectSource(source) rt := newClientSessionRuntime(conn, stopCtx, stopFn, queue, epoch) c.setClientSessionRuntimeWithCloseOld(rt, true) @@ -341,7 +340,7 @@ func (c *ClientCommon) startClientTransportRuntime(rt *clientSessionRuntime) err if c.useHeartBeat { go c.heartbeatLoop(transportStopCtx, rt.epoch) } - go c.readMessageLoop(transportStopCtx, rt.conn, rt.queue, rt.epoch) + go c.readMessageLoopAtRoute(transportStopCtx, clientSessionRouteFromRuntime(rt)) go c.loadMessageLoop(rt) return nil } @@ -440,10 +439,28 @@ func (c *ClientCommon) readMessage() { } func (c *ClientCommon) readMessageLoop(stopCtx context.Context, conn net.Conn, queue *stario.StarQueue, epoch uint64) { + route := c.clientSessionRouteSnapshot() + if route.binding == nil || route.binding.connSnapshot() != conn || route.binding.queueSnapshot() != queue || route.epoch != epoch { + route = clientSessionRoute{ + binding: newTransportBinding(conn, queue), + epoch: epoch, + sessionStopCtx: stopCtx, + transportStopCtx: stopCtx, + } + } + c.readMessageLoopAtRoute(stopCtx, route) +} + +func (c *ClientCommon) readMessageLoopAtRoute(stopCtx context.Context, route clientSessionRoute) { if stopCtx == nil { return } - binding := newTransportBinding(conn, queue) + binding := route.binding + if binding == nil { + return + } + conn := binding.connSnapshot() + queue := binding.queueSnapshot() dispatcher := c.clientInboundDispatcherSnapshot() if conn != nil && queue != nil && !isPacketTransportConn(conn) { reader := newTransportFrameReader(conn, queue) @@ -455,7 +472,7 @@ func (c *ClientCommon) readMessageLoop(stopCtx context.Context, conn net.Conn, q default: } payload, release, err := c.readTransportPayloadPooled(conn, reader) - if !c.handleTransportPayloadReadResultWithSession(stopCtx, binding, payload, release, err, epoch, dispatcher) { + if !c.handleTransportPayloadReadResultAtRoute(stopCtx, route, payload, release, err, dispatcher) { return } } @@ -469,7 +486,7 @@ func (c *ClientCommon) readMessageLoop(stopCtx context.Context, conn net.Conn, q default: } readNum, data, err := c.readFromTransportBindingWithBuffer(binding, buf) - if !c.handleTransportReadResultWithSessionDispatcher(stopCtx, conn, queue, readNum, data, err, epoch, dispatcher) { + if !c.handleTransportReadResultAtRoute(stopCtx, route, readNum, data, err, dispatcher) { return } } @@ -512,6 +529,7 @@ func (c *ClientCommon) loadMessageLoop(rt *clientSessionRuntime) { return } dispatcher := rt.inboundDispatcher + route := clientSessionRouteFromRuntime(rt) if dispatcher == nil { dispatcher = newInboundDispatcher() defer dispatcher.CloseAndWait() @@ -534,10 +552,10 @@ func (c *ClientCommon) loadMessageLoop(rt *clientSessionRuntime) { } msg := data c.wg.Add(1) - if !dispatcher.Dispatch(clientInboundDispatchSource(), func() { + if !dispatcher.DispatchSized(clientInboundDispatchSource(), len(msg.Msg), func() { defer c.wg.Done() now := time.Now() - if err := c.dispatchInboundTransportPayload(msg.Msg, now); err != nil { + if err := c.dispatchInboundTransportPayloadAtRoute(route, msg.Msg, now); err != nil { if c.showError || c.debugMode { fmt.Println("client decode envelope error", err) } diff --git a/client_send.go b/client_send.go index ebdf117..85679ea 100644 --- a/client_send.go +++ b/client_send.go @@ -17,7 +17,11 @@ func (c *ClientCommon) sendWithContext(ctx context.Context, msg TransferMsg) (Wa } func (c *ClientCommon) sendWithContextTimeout(ctx context.Context, msg TransferMsg, writeTimeout time.Duration) (WaitMsg, error) { - if err := c.ensureClientSendReady(); err != nil { + return c.sendWithContextTimeoutAtRoute(ctx, c.clientSessionRouteSnapshot(), msg, writeTimeout) +} + +func (c *ClientCommon) sendWithContextTimeoutAtRoute(ctx context.Context, route clientSessionRoute, msg TransferMsg, writeTimeout time.Duration) (WaitMsg, error) { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { return WaitMsg{}, err } if ctx == nil { @@ -36,7 +40,7 @@ func (c *ClientCommon) sendWithContextTimeout(ctx context.Context, msg TransferM if requiresSignalReplyWait(msg) { wait = c.getPendingWaitPool().createAndStore(msg) } - err = c.sendSignalEnvelopeMaybeReliable(env, msg) + err = c.sendSignalEnvelopeMaybeReliableAtRoute(route, env, msg) if err != nil { if requiresSignalReplyWait(msg) { c.getPendingWaitPool().removeAndClose(msg.ID) @@ -47,7 +51,11 @@ func (c *ClientCommon) sendWithContextTimeout(ctx context.Context, msg TransferM } func (c *ClientCommon) sendEnvelope(env Envelope) error { - if err := c.ensureClientSendReady(); err != nil { + return c.sendEnvelopeAtRoute(c.clientSessionRouteSnapshot(), env) +} + +func (c *ClientCommon) sendEnvelopeAtRoute(route clientSessionRoute, env Envelope) error { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { return err } payload, err := c.encodeEnvelopePayload(env) @@ -55,19 +63,23 @@ func (c *ClientCommon) sendEnvelope(env Envelope) error { return err } if batchedControlEnvelope(env) { - return c.writeControlPayloadToTransportTimeout(env.controlContext(), payload, env.controlPriority, env.controlTimeout) + return c.writeControlPayloadToTransportBindingTimeout(env.controlContext(), route.binding, payload, env.controlPriority, env.controlTimeout) } - return c.writePayloadToTransportContextTimeout(env.controlContext(), payload, env.controlTimeout) + return c.writePayloadToTransportBindingContextTimeout(env.controlContext(), route.binding, payload, env.controlTimeout) } func (c *ClientCommon) dispatchEnvelope(env Envelope, now time.Time) { + c.dispatchEnvelopeAtRoute(c.clientSessionRouteSnapshot(), env, now) +} + +func (c *ClientCommon) dispatchEnvelopeAtRoute(route clientSessionRoute, env Envelope, now time.Time) { switch env.Kind { case EnvelopeSignalAck: if c.handleSignalAckEnvelope(env) { return } case EnvelopeStreamData: - c.dispatchStreamEnvelope(env) + c.dispatchStreamEnvelopeAtRoute(route, env) return case EnvelopeSignal: transfer, err := unwrapTransferMsgEnvelope(env, c.sequenceDe) @@ -77,11 +89,12 @@ func (c *ClientCommon) dispatchEnvelope(env Envelope, now time.Time) { } return } - if c.handleReceivedSignalReliability(transfer) { + if c.handleReceivedSignalReliabilityAtRoute(route, transfer) { return } message := Message{ ServerConn: c, + clientRoute: route, TransferMsg: transfer, NetType: NET_CLIENT, Time: now, @@ -136,20 +149,31 @@ func (c *ClientCommon) sendWait(msg TransferMsg, timeout time.Duration) (Message } func (c *ClientCommon) sendCtx(msg TransferMsg, ctx context.Context) (Message, error) { + return c.sendCtxAtRoute(c.clientSessionRouteSnapshot(), msg, ctx) +} + +func (c *ClientCommon) sendCtxAtRoute(route clientSessionRoute, msg TransferMsg, ctx context.Context) (Message, error) { if ctx == nil { ctx = context.Background() } - data, err := c.sendWithContext(ctx, msg) + data, err := c.sendWithContextTimeoutAtRoute(ctx, route, msg, 0) if err != nil { return Message{}, publicContextSendError(ctx, err) } - stopCh := sessionStopChan(c.clientStopContextSnapshot()) + stopCh := sessionStopChan(route.sessionStopCtx) + transportStopCh := sessionStopChan(route.transportStopCtx) select { case <-ctx.Done(): c.getPendingWaitPool().removeAndClose(data.TransferMsg.ID) return Message{}, normalizeStreamDeadlineError(ctx.Err()) case <-stopCh: return Message{}, errServiceShutdown + case <-transportStopCh: + c.getPendingWaitPool().removeAndClose(data.TransferMsg.ID) + if route.sessionStopCtx != nil && route.sessionStopCtx.Err() != nil { + return Message{}, errServiceShutdown + } + return Message{}, transportDetachedSessionEpochError() case msg, ok := <-data.Reply: if !ok { return msg, pendingWaitClosedErrorWith(stopCh, clientTransportDetachedError(c)) @@ -159,11 +183,15 @@ func (c *ClientCommon) sendCtx(msg TransferMsg, ctx context.Context) (Message, e } func (c *ClientCommon) SendObjCtx(ctx context.Context, key string, val interface{}) (Message, error) { + return c.sendObjCtxAtRoute(ctx, c.clientSessionRouteSnapshot(), key, val) +} + +func (c *ClientCommon) sendObjCtxAtRoute(ctx context.Context, route clientSessionRoute, key string, val interface{}) (Message, error) { data, err := c.sequenceEn(val) if err != nil { return Message{}, err } - return c.sendCtx(TransferMsg{ + return c.sendCtxAtRoute(route, TransferMsg{ Key: key, Value: data, Type: MSG_SYNC_ASK, diff --git a/client_session_route.go b/client_session_route.go new file mode 100644 index 0000000..6b3a35b --- /dev/null +++ b/client_session_route.go @@ -0,0 +1,108 @@ +package notify + +import "context" + +// clientSessionRoute pins a send or inbound dispatch to one physical client +// transport. The logical client session can survive a transport reattach, so +// the epoch alone is not sufficient to prevent an old operation from crossing +// onto the replacement connection. +type clientSessionRoute struct { + runtime *clientSessionRuntime + binding *transportBinding + epoch uint64 + sessionStopCtx context.Context + transportStopCtx context.Context +} + +func clientSessionRouteFromRuntime(rt *clientSessionRuntime) clientSessionRoute { + if rt == nil { + return clientSessionRoute{} + } + transportStopCtx := rt.transportStopCtx + if transportStopCtx == nil { + transportStopCtx = rt.stopCtx + } + return clientSessionRoute{ + runtime: rt, + binding: rt.transport, + epoch: rt.epoch, + sessionStopCtx: rt.stopCtx, + transportStopCtx: transportStopCtx, + } +} + +func clientSessionRouteContext(route clientSessionRoute) context.Context { + if route.transportStopCtx != nil { + return route.transportStopCtx + } + return route.sessionStopCtx +} + +func sameClientSessionRoute(left, right clientSessionRoute) bool { + if left.runtime != nil || right.runtime != nil { + return left.runtime != nil && left.runtime == right.runtime + } + if left.binding != nil || right.binding != nil { + return left.binding != nil && left.binding == right.binding && + (left.epoch == 0 || right.epoch == 0 || left.epoch == right.epoch) + } + return left.epoch != 0 && left.epoch == right.epoch +} + +func (r clientSessionRoute) bound() bool { + return r.binding != nil || r.epoch != 0 +} + +func (r clientSessionRoute) inboundQueueKey() interface{} { + if r.binding != nil { + return r.binding + } + return "b612" +} + +func (c *ClientCommon) clientSessionRouteSnapshot() clientSessionRoute { + if c == nil { + return clientSessionRoute{} + } + return clientSessionRouteFromRuntime(c.clientSessionRuntimeSnapshot()) +} + +func (c *ClientCommon) clientSessionRouteCurrent(route clientSessionRoute) bool { + if c == nil { + return false + } + if !route.bound() { + return true + } + current := c.clientSessionRuntimeSnapshot() + if current == nil { + return false + } + if route.epoch != 0 && current.epoch != route.epoch { + return false + } + return route.binding != nil && current.transport == route.binding +} + +func (c *ClientCommon) ensureClientSessionRouteSendReady(route clientSessionRoute) error { + if err := c.ensureClientSendReady(); err != nil { + return err + } + if !c.clientSessionRouteCurrent(route) { + return transportDetachedSessionEpochError() + } + if route.transportStopCtx != nil { + select { + case <-route.transportStopCtx.Done(): + if route.sessionStopCtx != nil && route.sessionStopCtx.Err() != nil { + return errServiceShutdown + } + return transportDetachedSessionEpochError() + default: + } + } + if route.binding == nil { + return clientTransportDetachedError(c) + } + return nil +} diff --git a/client_session_runtime.go b/client_session_runtime.go index 35dd7da..e85669e 100644 --- a/client_session_runtime.go +++ b/client_session_runtime.go @@ -196,6 +196,16 @@ func (c *ClientCommon) attachClientSessionTransport(conn net.Conn) error { if rt.transportStopFn != nil { rt.transportStopFn() } + oldRoute := clientSessionRouteFromRuntime(rt) + if streamRuntime := c.getStreamRuntime(); streamRuntime != nil { + streamRuntime.closeClientRoute(oldRoute, errTransportDetached) + } + if bulkRuntime := c.getBulkRuntime(); bulkRuntime != nil { + bulkRuntime.closeClientRoute(oldRoute, errTransportDetached) + } + // A sidecar is physically tied to the old primary transport. Retire it + // before publishing the replacement route so a new bulk cannot reuse it. + c.closeClientDedicatedSidecarWithError(errTransportDetached) next := *rt next.transport = newTransportBinding(conn, rt.queue) next.transportAttached = true diff --git a/client_session_runtime_test.go b/client_session_runtime_test.go index 027ba14..db2751c 100644 --- a/client_session_runtime_test.go +++ b/client_session_runtime_test.go @@ -3,6 +3,7 @@ package notify import ( "b612.me/stario" "context" + "errors" "io" "math" "net" @@ -345,6 +346,76 @@ func TestAttachClientSessionTransportRebindsRuntimeAndDispatchesOnNewConn(t *tes } } +func TestAttachClientSessionTransportRetiresOldRouteTransfersAndSidecars(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) + oldLeft, oldRight := net.Pipe() + defer oldRight.Close() + client.setClientSessionRuntime(newClientSessionRuntime(oldLeft, stopCtx, stopFn, queue, 17)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + oldRoute := client.clientSessionRouteSnapshot() + + streamRuntime := client.getStreamRuntime() + stream := newStreamHandle(stopCtx, streamRuntime, clientFileScope(), StreamOpenRequest{ + StreamID: "old-route-stream", + DataID: 1, + }, oldRoute.epoch, nil, nil, 0, nil, nil, nil, streamRuntime.configSnapshot()) + stream.setClientSessionRoute(oldRoute) + if err := streamRuntime.register(clientFileScope(), stream); err != nil { + t.Fatalf("register old-route stream: %v", err) + } + defer stream.markReset(io.ErrClosedPipe) + + bulkRuntime := client.getBulkRuntime() + bulk := newBulkHandle(stopCtx, bulkRuntime, clientFileScope(), BulkOpenRequest{ + BulkID: "old-route-bulk", + DataID: 1, + }, oldRoute.epoch, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.setClientSessionRoute(oldRoute) + if err := bulkRuntime.registerOutbound(clientFileScope(), bulk); err != nil { + t.Fatalf("register old-route bulk: %v", err) + } + defer bulk.markReset(io.ErrClosedPipe) + + sidecarLeft, sidecarRight := net.Pipe() + defer sidecarRight.Close() + sidecar := newBulkDedicatedSidecar(sidecarLeft, 1) + if active, installed := client.installClientDedicatedSidecar(1, sidecar); !installed || active != sidecar { + t.Fatalf("install old-route sidecar = %p/%v, want %p/true", active, installed, sidecar) + } + + newLeft, newRight := net.Pipe() + defer newRight.Close() + if err := client.attachClientSessionTransport(newLeft); err != nil { + t.Fatalf("attach replacement client transport: %v", err) + } + + if err := stream.resetErrSnapshot(); !errors.Is(err, errTransportDetached) { + t.Fatalf("old-route stream reset error = %v, want transport detached", err) + } + if _, ok := streamRuntime.lookup(clientFileScope(), stream.ID()); ok { + t.Fatal("old-route stream remained registered after transport reattach") + } + if err := bulk.resetErrSnapshot(); !errors.Is(err, errTransportDetached) { + t.Fatalf("old-route bulk reset error = %v, want transport detached", err) + } + if _, ok := bulkRuntime.lookup(clientFileScope(), bulk.ID()); ok { + t.Fatal("old-route bulk remained registered after transport reattach") + } + if got := client.clientDedicatedSidecarSnapshotForLane(1); got != nil { + t.Fatalf("old-route dedicated sidecar remained installed: %p", got) + } + _ = sidecarRight.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := sidecarRight.Read(make([]byte, 1)); err == nil { + t.Fatal("old-route dedicated sidecar connection remained open") + } +} + func TestSetClientSessionRuntimeStopsOldBindingWorkersOnReattach(t *testing.T) { client := NewClient().(*ClientCommon) diff --git a/client_stream.go b/client_stream.go index 906b942..3bba144 100644 --- a/client_stream.go +++ b/client_stream.go @@ -21,44 +21,72 @@ func (c *ClientCommon) OpenStream(ctx context.Context, opt StreamOpenOptions) (S if runtime == nil { return nil, errStreamRuntimeNil } + route := c.clientSessionRouteSnapshot() + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return nil, err + } + scope := clientFileScope() req := clientStreamRequest(runtime, opt) if req.StreamID == "" { return nil, errStreamIDEmpty } - if _, exists := runtime.lookup(clientFileScope(), req.StreamID); exists { - return nil, errStreamAlreadyExists + if existing, exists := runtime.lookup(scope, req.StreamID); exists { + if existing.acceptsClientSessionRoute(route) { + return nil, errStreamAlreadyExists + } + existing.markReset(transportDetachedSessionEpochError()) } - resp, err := sendStreamOpenClient(ctx, c, req) + dataID, err := runtime.reserveDataID(scope) if err != nil { return nil, err } - if resp.DataID != 0 { - req.DataID = resp.DataID + req.DataID = dataID + parent := clientSessionRouteContext(route) + if parent == nil { + parent = c.clientStopContextSnapshot() } - if resp.FastPathVersion != 0 { - req.FastPathVersion = resp.FastPathVersion - } else { - req.FastPathVersion = streamFastPathVersionV1 - } - req.Metadata = mergeStreamMetadata(req.Metadata, resp.Metadata) - stream := newStreamHandle(c.clientStopContextSnapshot(), runtime, clientFileScope(), req, c.currentClientSessionEpoch(), nil, nil, resp.TransportGeneration, clientStreamCloseSender(c), clientStreamResetSender(c), clientStreamDataSender(c, c.currentClientSessionEpoch()), runtime.configSnapshot()) + stream := newStreamHandle(parent, runtime, scope, req, route.epoch, nil, nil, 0, clientStreamCloseSender(c), clientStreamResetSender(c), clientStreamDataSender(c, route), runtime.configSnapshot()) stream.setClientSnapshotOwner(c) - stream.setAddrSnapshot(c.clientStreamAddrSnapshot()) - if err := runtime.register(clientFileScope(), stream); err != nil { - _, _ = sendStreamResetClient(context.Background(), c, StreamResetRequest{ - StreamID: req.StreamID, - Error: err.Error(), - }) + stream.setClientSessionRoute(route) + stream.setAddrSnapshot(c.clientStreamAddrSnapshotAtRoute(route)) + if err := runtime.adoptReserved(scope, stream); err != nil { + runtime.releaseDataID(scope, req.DataID) return nil, err } + resp, err := sendStreamOpenClientAtRoute(ctx, c, route, req) + if err != nil { + c.bestEffortStreamResetAtRoute(route, StreamResetRequest{StreamID: req.StreamID, DataID: req.DataID, Error: err.Error()}) + stream.markReset(err) + return nil, err + } + if resp.DataID != 0 && resp.DataID != req.DataID { + err = errStreamAlreadyExists + c.bestEffortStreamResetAtRoute(route, StreamResetRequest{StreamID: req.StreamID, Error: "stream data id mismatch"}) + stream.markReset(err) + return nil, err + } + if resp.FastPathVersion != 0 { + stream.setFastPathVersion(resp.FastPathVersion) + } else { + stream.setFastPathVersion(streamFastPathVersionV1) + } + stream.metadata = mergeStreamMetadata(req.Metadata, resp.Metadata) + stream.setTransportGeneration(resp.TransportGeneration) return stream, nil } func (c *ClientCommon) clientStreamAddrSnapshot() (net.Addr, net.Addr) { + return c.clientStreamAddrSnapshotAtRoute(c.clientSessionRouteSnapshot()) +} + +func (c *ClientCommon) clientStreamAddrSnapshotAtRoute(route clientSessionRoute) (net.Addr, net.Addr) { if c == nil { return nil, nil } - conn := c.clientTransportConnSnapshot() + var conn net.Conn + if route.binding != nil { + conn = route.binding.connSnapshot() + } if conn == nil { return nil, nil } @@ -82,8 +110,9 @@ func clientStreamRequest(runtime *streamRuntime, opt StreamOpenOptions) StreamOp func clientStreamCloseSender(c *ClientCommon) streamCloseSender { return func(ctx context.Context, stream *streamHandle, full bool) error { - _, err := sendStreamCloseClient(ctx, c, StreamCloseRequest{ + _, err := sendStreamCloseClientAtRoute(ctx, c, stream.clientSessionRouteSnapshot(), StreamCloseRequest{ StreamID: stream.ID(), + DataID: stream.dataIDSnapshot(), Full: full, }) return err @@ -92,21 +121,23 @@ func clientStreamCloseSender(c *ClientCommon) streamCloseSender { func clientStreamResetSender(c *ClientCommon) streamResetSender { return func(ctx context.Context, stream *streamHandle, message string) error { - _, err := sendStreamResetClient(ctx, c, StreamResetRequest{ - StreamID: stream.ID(), - Error: message, + _, err := sendStreamResetClientAtRoute(ctx, c, stream.clientSessionRouteSnapshot(), StreamResetRequest{ + StreamID: stream.ID(), + DataID: stream.dataIDSnapshot(), + Error: message, + RecordFailure: stream.recordResetFailure(), }) return err } } -func clientStreamDataSender(c *ClientCommon, epoch uint64) streamDataSender { +func clientStreamDataSender(c *ClientCommon, route clientSessionRoute) streamDataSender { return func(ctx context.Context, stream *streamHandle, chunk []byte) error { if c == nil { return errStreamClientNil } - if epoch != 0 && !c.isClientSessionEpochCurrent(epoch) { - return errTransportDetached + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return err } if ctx != nil { select { @@ -116,8 +147,17 @@ func clientStreamDataSender(c *ClientCommon, epoch uint64) streamDataSender { } } if dataID := stream.dataIDSnapshot(); dataID != 0 { - return c.sendFastStreamData(ctx, stream, chunk) + return c.sendFastStreamDataAtRoute(ctx, route, stream, chunk) } - return c.sendEnvelope(newStreamDataEnvelope(stream.ID(), chunk)) + return c.sendEnvelopeAtRoute(route, newStreamDataEnvelope(stream.ID(), chunk)) } } + +func (c *ClientCommon) bestEffortStreamResetAtRoute(route clientSessionRoute, req StreamResetRequest) { + if c == nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), streamDispatchRejectTimeout) + defer cancel() + _, _ = sendStreamResetClientAtRoute(ctx, c, route, req) +} diff --git a/client_transport.go b/client_transport.go index 86c9f52..85fc504 100644 --- a/client_transport.go +++ b/client_transport.go @@ -100,31 +100,43 @@ func (c *ClientCommon) handleTransportReadResultWithSession(stopCtx context.Cont } func (c *ClientCommon) handleTransportReadResultWithSessionDispatcher(stopCtx context.Context, conn net.Conn, queue *stario.StarQueue, readNum int, data []byte, err error, epoch uint64, dispatcher *inboundDispatcher) bool { - binding := newTransportBinding(conn, queue) + route := c.clientSessionRouteSnapshot() + if route.binding == nil || route.binding.connSnapshot() != conn || route.binding.queueSnapshot() != queue || route.epoch != epoch { + route = clientSessionRoute{binding: newTransportBinding(conn, queue), epoch: epoch, sessionStopCtx: stopCtx, transportStopCtx: stopCtx} + } + return c.handleTransportReadResultAtRoute(stopCtx, route, readNum, data, err, dispatcher) +} + +func (c *ClientCommon) handleTransportReadResultAtRoute(stopCtx context.Context, route clientSessionRoute, readNum int, data []byte, err error, dispatcher *inboundDispatcher) bool { + binding := route.binding + queue := binding.queueSnapshot() if err == os.ErrDeadlineExceeded { if readNum != 0 && queue != nil { - if !c.pushMessageFast(queue, data[:readNum], dispatcher) { - queue.ParseMessage(data[:readNum], "b612") + if !c.pushMessageFastAtRoute(route, queue, data[:readNum], dispatcher) { + queue.ParseMessage(data[:readNum], route.inboundQueueKey()) } } return true } if err != nil { - if c.showError || c.debugMode { - fmt.Println("client read error", err) - } select { case <-sessionStopChan(stopCtx): c.closeClientTransportBinding(binding) return false default: } - c.stopClientSessionIfCurrent(epoch, "client read error", err) + // An expected shutdown closes the socket, which surfaces here as a + // read on a closed connection. Only report reads that happen while the + // session is still supposed to be running. + if c.showError || c.debugMode { + fmt.Println("client read error", err) + } + c.stopClientSessionIfCurrent(route.epoch, "client read error", err) return false } if queue != nil { - if !c.pushMessageFast(queue, data[:readNum], dispatcher) { - queue.ParseMessage(data[:readNum], "b612") + if !c.pushMessageFastAtRoute(route, queue, data[:readNum], dispatcher) { + queue.ParseMessage(data[:readNum], route.inboundQueueKey()) } } return true @@ -144,6 +156,15 @@ func (c *ClientCommon) readTransportPayloadPooled(conn net.Conn, reader *stario. } func (c *ClientCommon) handleTransportPayloadReadResultWithSession(stopCtx context.Context, binding *transportBinding, payload []byte, release func(), err error, epoch uint64, dispatcher *inboundDispatcher) bool { + route := c.clientSessionRouteSnapshot() + if route.binding != binding || route.epoch != epoch { + route = clientSessionRoute{binding: binding, epoch: epoch, sessionStopCtx: stopCtx, transportStopCtx: stopCtx} + } + return c.handleTransportPayloadReadResultAtRoute(stopCtx, route, payload, release, err, dispatcher) +} + +func (c *ClientCommon) handleTransportPayloadReadResultAtRoute(stopCtx context.Context, route clientSessionRoute, payload []byte, release func(), err error, dispatcher *inboundDispatcher) bool { + binding := route.binding if err == os.ErrDeadlineExceeded { return true } @@ -151,23 +172,28 @@ func (c *ClientCommon) handleTransportPayloadReadResultWithSession(stopCtx conte if release != nil { release() } - if c.showError || c.debugMode { - fmt.Println("client read error", err) - } select { case <-sessionStopChan(stopCtx): c.closeClientTransportBinding(binding) return false default: } - c.stopClientSessionIfCurrent(epoch, "client read error", err) + // See handleTransportReadResultAtRoute: shutdown reads are expected. + if c.showError || c.debugMode { + fmt.Println("client read error", err) + } + c.stopClientSessionIfCurrent(route.epoch, "client read error", err) return false } - c.dispatchTransportPayloadFast(payload, release, dispatcher) + c.dispatchTransportPayloadFastAtRoute(route, payload, release, dispatcher) return true } func (c *ClientCommon) dispatchTransportPayloadFast(payload []byte, release func(), dispatcher *inboundDispatcher) { + c.dispatchTransportPayloadFastAtRoute(c.clientSessionRouteSnapshot(), payload, release, dispatcher) +} + +func (c *ClientCommon) dispatchTransportPayloadFastAtRoute(route clientSessionRoute, payload []byte, release func(), dispatcher *inboundDispatcher) { if len(payload) == 0 { if release != nil { release() @@ -181,12 +207,12 @@ func (c *ClientCommon) dispatchTransportPayloadFast(payload []byte, release func } return } - if c.tryDispatchBorrowedTransportPlain(plain, plainRelease) { + if c.tryDispatchBorrowedTransportPlainAtRoute(route, plain, plainRelease) { return } if dispatcher == nil { now := time.Now() - err := c.dispatchInboundTransportPlain(plain, now) + err := c.dispatchInboundTransportPlainAtRoute(route, plain, now) if plainRelease != nil { plainRelease() } @@ -201,10 +227,10 @@ func (c *ClientCommon) dispatchTransportPayloadFast(payload []byte, release func plainRelease() } c.wg.Add(1) - if !dispatcher.Dispatch(clientInboundDispatchSource(), func() { + if !dispatcher.DispatchSized(clientInboundDispatchSource(), len(owned), func() { defer c.wg.Done() now := time.Now() - if err := c.dispatchInboundTransportPlain(owned, now); err != nil && (c.showError || c.debugMode) { + if err := c.dispatchInboundTransportPlainAtRoute(route, owned, now); err != nil && (c.showError || c.debugMode) { fmt.Println("client decode envelope error", err) } }) { @@ -213,11 +239,15 @@ func (c *ClientCommon) dispatchTransportPayloadFast(payload []byte, release func } func (c *ClientCommon) pushMessageFast(queue *stario.StarQueue, data []byte, dispatcher *inboundDispatcher) bool { + return c.pushMessageFastAtRoute(c.clientSessionRouteSnapshot(), queue, data, dispatcher) +} + +func (c *ClientCommon) pushMessageFastAtRoute(route clientSessionRoute, queue *stario.StarQueue, data []byte, dispatcher *inboundDispatcher) bool { if queue == nil || dispatcher == nil || len(data) == 0 { return false } - if err := queue.ParseMessageView(data, "b612", func(frame stario.FrameView) error { - c.dispatchTransportPayloadFast(frame.Payload, nil, dispatcher) + if err := queue.ParseMessageView(data, route.inboundQueueKey(), func(frame stario.FrameView) error { + c.dispatchTransportPayloadFastAtRoute(route, frame.Payload, nil, dispatcher) return nil }); err != nil && (c.showError || c.debugMode) { fmt.Println("client parse inbound frame error", err) @@ -248,6 +278,10 @@ func (c *ClientCommon) writePayloadToTransportContext(ctx context.Context, paylo func (c *ClientCommon) writePayloadToTransportContextTimeout(ctx context.Context, payload []byte, writeTimeout time.Duration) error { binding := c.clientTransportBindingSnapshot() + return c.writePayloadToTransportBindingContextTimeout(ctx, binding, payload, writeTimeout) +} + +func (c *ClientCommon) writePayloadToTransportBindingContextTimeout(ctx context.Context, binding *transportBinding, payload []byte, writeTimeout time.Duration) error { if binding == nil { return net.ErrClosed } @@ -273,6 +307,10 @@ func (c *ClientCommon) writeControlPayloadToTransport(ctx context.Context, paylo func (c *ClientCommon) writeControlPayloadToTransportTimeout(ctx context.Context, payload []byte, priority controlPriority, writeTimeout time.Duration) error { binding := c.clientTransportBindingSnapshot() + return c.writeControlPayloadToTransportBindingTimeout(ctx, binding, payload, priority, writeTimeout) +} + +func (c *ClientCommon) writeControlPayloadToTransportBindingTimeout(ctx context.Context, binding *transportBinding, payload []byte, priority controlPriority, writeTimeout time.Duration) error { if binding == nil { return net.ErrClosed } @@ -282,11 +320,11 @@ func (c *ClientCommon) writeControlPayloadToTransportTimeout(ctx context.Context } conn := binding.connSnapshot() if conn == nil || isPacketTransportConn(conn) { - return c.writePayloadToTransportContextTimeout(ctx, payload, writeTimeout) + return c.writePayloadToTransportBindingContextTimeout(ctx, binding, payload, writeTimeout) } sender := binding.controlBatchSenderSnapshot() if sender == nil { - return c.writePayloadToTransportContextTimeout(ctx, payload, writeTimeout) + return c.writePayloadToTransportBindingContextTimeout(ctx, binding, payload, writeTimeout) } return sender.submitContext(ctx, payload, shorterPositiveDuration(c.maxWriteTimeoutSnapshot(), writeTimeout), priority) } diff --git a/control_batch_sender.go b/control_batch_sender.go index d81eb97..6204cd7 100644 --- a/control_batch_sender.go +++ b/control_batch_sender.go @@ -611,12 +611,25 @@ func (s *controlBatchSender) controlBatchWaitContext(requests []controlBatchRequ base = s.stopCtx } ctx, cancel := context.WithCancel(base) - stops := make([]func() bool, 0, len(requests)) + var remaining atomic.Int32 for _, item := range requests { - if item.ctx == nil || item.ctx.Done() == nil { - continue + if item.ctx != nil && item.ctx.Done() != nil { + remaining.Add(1) + } + } + stops := make([]func() bool, 0, len(requests)) + if remaining.Load() > 0 { + cancelWhenAllDone := func() { + if remaining.Add(-1) == 0 { + cancel() + } + } + for _, item := range requests { + if item.ctx == nil || item.ctx.Done() == nil { + continue + } + stops = append(stops, context.AfterFunc(item.ctx, cancelWhenAllDone)) } - stops = append(stops, context.AfterFunc(item.ctx, cancel)) } return ctx, func() { for _, stop := range stops { diff --git a/file_receiver_test.go b/file_receiver_test.go index fb68038..f48ec8f 100644 --- a/file_receiver_test.go +++ b/file_receiver_test.go @@ -6,6 +6,7 @@ import ( "encoding/hex" "os" "path/filepath" + "runtime" "testing" "time" ) @@ -257,7 +258,12 @@ func TestFileReceivePoolAppliesMetaModeAndModTime(t *testing.T) { if err != nil { t.Fatalf("Stat failed: %v", err) } - if got, want := info.Mode().Perm(), wantMode; got != want { + wantPerm := wantMode + if runtime.GOOS == "windows" { + // Windows Chmod only controls the owner write bit (read-only attribute). + wantPerm = 0o666 + } + if got, want := info.Mode().Perm(), wantPerm; got != want { t.Fatalf("mode mismatch: got %o want %o", got, want) } gotMTime := info.ModTime().Truncate(time.Second) diff --git a/inbound_dispatcher.go b/inbound_dispatcher.go index 1e32318..e271303 100644 --- a/inbound_dispatcher.go +++ b/inbound_dispatcher.go @@ -8,53 +8,170 @@ import ( const defaultInboundDispatchSource = "_notify.default_inbound_source" +// Inbound dispatch runs one serial worker per source so messages from the same +// connection keep their relative order. The pending queue is bounded both per +// source and in total: a slow handler can never let one connection grow the +// queue without limit, and a saturated connection can never starve the others. +// Callers block until there is room, which pushes backpressure to the transport +// reader instead of growing memory; CloseAndWait unblocks every caller. +const ( + defaultInboundDispatchQueueLimit = 4096 + defaultInboundDispatchQueueBytes = 64 << 20 + defaultInboundSourceQueueLimit = 512 + defaultInboundSourceQueueBytes = 16 << 20 +) + +type inboundDispatchItem struct { + size int + fn func() +} + type inboundDispatcher struct { - mu sync.Mutex - closed bool - workers map[string]*inboundDispatchWorker - wg sync.WaitGroup + mu sync.Mutex + closed bool + closeCh chan struct{} + roomCh chan struct{} + queued int + queuedBytes int + maxItems int + maxBytes int + sourceItems int + sourceBytes int + workers map[string]*inboundDispatchWorker + wg sync.WaitGroup } type inboundDispatchWorker struct { - queue []func() - running bool + queue []inboundDispatchItem + running bool + queued int + queuedBytes int } func newInboundDispatcher() *inboundDispatcher { + return newInboundDispatcherWithCaps( + defaultInboundDispatchQueueLimit, + defaultInboundDispatchQueueBytes, + defaultInboundSourceQueueLimit, + defaultInboundSourceQueueBytes, + ) +} + +// newInboundDispatcherWithLimits sizes a dispatcher for a single source: the +// per-source caps equal the global caps. +func newInboundDispatcherWithLimits(maxItems int, maxBytes int) *inboundDispatcher { + return newInboundDispatcherWithCaps(maxItems, maxBytes, maxItems, maxBytes) +} + +func newInboundDispatcherWithCaps(maxItems int, maxBytes int, sourceItems int, sourceBytes int) *inboundDispatcher { + if maxItems <= 0 { + maxItems = defaultInboundDispatchQueueLimit + } + if maxBytes <= 0 { + maxBytes = defaultInboundDispatchQueueBytes + } + if sourceItems <= 0 || sourceItems > maxItems { + sourceItems = maxItems + } + if sourceBytes <= 0 || sourceBytes > maxBytes { + sourceBytes = maxBytes + } return &inboundDispatcher{ - workers: make(map[string]*inboundDispatchWorker), + closeCh: make(chan struct{}), + roomCh: make(chan struct{}, 1), + maxItems: maxItems, + maxBytes: maxBytes, + sourceItems: sourceItems, + sourceBytes: sourceBytes, + workers: make(map[string]*inboundDispatchWorker), } } +// Dispatch queues fn for the given source without byte accounting. Prefer +// DispatchSized when the queued payload size is known. Like DispatchSized it +// blocks while the queue is at its limit. func (d *inboundDispatcher) Dispatch(source string, fn func()) bool { + return d.DispatchSized(source, 0, fn) +} + +// DispatchSized queues fn for the given source. It blocks while the source or +// the dispatcher is at its item or byte limit, and returns false once the +// dispatcher is closed. A parked caller is released by CloseAndWait, so the +// owner of the reader must keep CloseAndWait reachable (a concurrent closer, or +// the reader's own stop path once the wait unblocks). +func (d *inboundDispatcher) DispatchSized(source string, size int, fn func()) bool { if d == nil || fn == nil { return false } if source == "" { source = defaultInboundDispatchSource } - d.mu.Lock() - if d.closed { + if size < 0 { + size = 0 + } + for { + d.mu.Lock() + if d.closed { + d.mu.Unlock() + return false + } + worker := d.workers[source] + if worker == nil { + worker = &inboundDispatchWorker{} + d.workers[source] = worker + } + if d.roomLocked(worker, size) { + worker.queue = append(worker.queue, inboundDispatchItem{size: size, fn: fn}) + worker.queued++ + worker.queuedBytes += size + d.queued++ + d.queuedBytes += size + if worker.running { + d.mu.Unlock() + return true + } + worker.running = true + d.wg.Add(1) + d.mu.Unlock() + go d.run(source, worker) + return true + } d.mu.Unlock() + // roomCh is signalled whenever a queued item is consumed; closeCh + // releases the caller during shutdown. + select { + case <-d.roomCh: + case <-d.closeCh: + return false + } + } +} + +func (d *inboundDispatcher) roomLocked(worker *inboundDispatchWorker, size int) bool { + if d.maxItems > 0 && d.queued >= d.maxItems { return false } - worker := d.workers[source] - if worker == nil { - worker = &inboundDispatchWorker{} - d.workers[source] = worker + if d.sourceItems > 0 && worker.queued >= d.sourceItems { + return false } - worker.queue = append(worker.queue, fn) - if worker.running { - d.mu.Unlock() - return true + // Always admit at least one item per scope so an oversized payload cannot + // deadlock the reader behind an empty queue. + if d.maxBytes > 0 && d.queued > 0 && d.queuedBytes+size > d.maxBytes { + return false + } + if d.sourceBytes > 0 && worker.queued > 0 && worker.queuedBytes+size > d.sourceBytes { + return false } - worker.running = true - d.wg.Add(1) - d.mu.Unlock() - go d.run(source, worker) return true } +func (d *inboundDispatcher) signalRoomLocked() { + select { + case d.roomCh <- struct{}{}: + default: + } +} + func (d *inboundDispatcher) run(source string, worker *inboundDispatchWorker) { defer d.wg.Done() for { @@ -64,14 +181,20 @@ func (d *inboundDispatcher) run(source string, worker *inboundDispatchWorker) { if current := d.workers[source]; current == worker { delete(d.workers, source) } + d.signalRoomLocked() d.mu.Unlock() return } - fn := worker.queue[0] - worker.queue[0] = nil + item := worker.queue[0] + worker.queue[0] = inboundDispatchItem{} worker.queue = worker.queue[1:] + worker.queued-- + worker.queuedBytes -= item.size + d.queued-- + d.queuedBytes -= item.size + d.signalRoomLocked() d.mu.Unlock() - fn() + item.fn() } } @@ -80,7 +203,11 @@ func (d *inboundDispatcher) CloseAndWait() { return } d.mu.Lock() - d.closed = true + if !d.closed { + d.closed = true + close(d.closeCh) + } + d.signalRoomLocked() d.mu.Unlock() d.wg.Wait() } diff --git a/inbound_dispatcher_test.go b/inbound_dispatcher_test.go index a0919d1..4a3227b 100644 --- a/inbound_dispatcher_test.go +++ b/inbound_dispatcher_test.go @@ -2,6 +2,7 @@ package notify import ( "sync" + "sync/atomic" "testing" "time" ) @@ -101,3 +102,233 @@ func indexOfString(list []string, target string) int { } return -1 } + +func TestInboundDispatcherBoundsQueuedItems(t *testing.T) { + dispatcher := newInboundDispatcherWithLimits(1, 1<<20) + defer dispatcher.CloseAndWait() + + releaseFirst := make(chan struct{}) + firstStarted := make(chan struct{}) + if !dispatcher.DispatchSized("alpha", 1, func() { + close(firstStarted) + <-releaseFirst + }) { + t.Fatal("dispatch first item failed") + } + select { + case <-firstStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first item") + } + + if !dispatcher.DispatchSized("alpha", 1, func() {}) { + t.Fatal("dispatch second item failed") + } + + thirdDone := make(chan bool, 1) + go func() { + thirdDone <- dispatcher.DispatchSized("alpha", 1, func() {}) + }() + select { + case <-thirdDone: + t.Fatal("dispatch exceeded the queue limit without backpressure") + case <-time.After(100 * time.Millisecond): + } + + close(releaseFirst) + select { + case ok := <-thirdDone: + if !ok { + t.Fatal("blocked dispatch failed after room became available") + } + case <-time.After(time.Second): + t.Fatal("blocked dispatch was not released when the queue drained") + } +} + +func TestInboundDispatcherBoundsQueuedBytes(t *testing.T) { + dispatcher := newInboundDispatcherWithLimits(1000, 10) + defer dispatcher.CloseAndWait() + + releaseFirst := make(chan struct{}) + firstStarted := make(chan struct{}) + if !dispatcher.DispatchSized("alpha", 1, func() { + close(firstStarted) + <-releaseFirst + }) { + t.Fatal("dispatch first item failed") + } + select { + case <-firstStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first item") + } + + if !dispatcher.DispatchSized("alpha", 6, func() {}) { + t.Fatal("dispatch second item failed") + } + + thirdDone := make(chan bool, 1) + go func() { + thirdDone <- dispatcher.DispatchSized("alpha", 6, func() {}) + }() + select { + case <-thirdDone: + t.Fatal("dispatch exceeded the byte limit without backpressure") + case <-time.After(100 * time.Millisecond): + } + + close(releaseFirst) + select { + case ok := <-thirdDone: + if !ok { + t.Fatal("blocked dispatch failed after room became available") + } + case <-time.After(time.Second): + t.Fatal("blocked dispatch was not released when the queue drained") + } +} + +func TestInboundDispatcherCloseUnblocksBlockedDispatch(t *testing.T) { + dispatcher := newInboundDispatcherWithLimits(1, 1<<20) + + releaseFirst := make(chan struct{}) + firstStarted := make(chan struct{}) + if !dispatcher.DispatchSized("alpha", 1, func() { + close(firstStarted) + <-releaseFirst + }) { + t.Fatal("dispatch first item failed") + } + select { + case <-firstStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first item") + } + if !dispatcher.DispatchSized("alpha", 1, func() {}) { + t.Fatal("dispatch second item failed") + } + + blockedDone := make(chan bool, 1) + go func() { + blockedDone <- dispatcher.DispatchSized("alpha", 1, func() {}) + }() + select { + case <-blockedDone: + t.Fatal("dispatch exceeded the queue limit without backpressure") + case <-time.After(100 * time.Millisecond): + } + + closeWaitDone := make(chan struct{}) + go func() { + dispatcher.CloseAndWait() + close(closeWaitDone) + }() + + select { + case ok := <-blockedDone: + if ok { + t.Fatal("dispatch accepted work after the dispatcher closed") + } + case <-time.After(time.Second): + t.Fatal("CloseAndWait did not release the blocked dispatch") + } + + close(releaseFirst) + select { + case <-closeWaitDone: + case <-time.After(time.Second): + t.Fatal("CloseAndWait did not return after in-flight work finished") + } +} + +func TestInboundDispatcherConcurrentDispatchDrainsEveryItem(t *testing.T) { + dispatcher := newInboundDispatcherWithLimits(8, 1<<20) + + const producers = 8 + const perProducer = 64 + var handled atomic.Int64 + var wg sync.WaitGroup + for p := 0; p < producers; p++ { + wg.Add(1) + go func() { + defer wg.Done() + for i := 0; i < perProducer; i++ { + if !dispatcher.DispatchSized("shared", 1, func() { + handled.Add(1) + time.Sleep(50 * time.Microsecond) + }) { + t.Errorf("dispatch rejected before close") + return + } + } + }() + } + wg.Wait() + dispatcher.CloseAndWait() + + if got, want := handled.Load(), int64(producers*perProducer); got != want { + t.Fatalf("handled=%d, want %d", got, want) + } + dispatcher.mu.Lock() + queued, queuedBytes := dispatcher.queued, dispatcher.queuedBytes + dispatcher.mu.Unlock() + if queued != 0 || queuedBytes != 0 { + t.Fatalf("queue accounting leaked: queued=%d bytes=%d", queued, queuedBytes) + } +} + +func TestInboundDispatcherPerSourceLimitDoesNotStarveOtherSources(t *testing.T) { + dispatcher := newInboundDispatcherWithCaps(64, 1<<20, 1, 1<<20) + defer dispatcher.CloseAndWait() + + releaseFirst := make(chan struct{}) + firstStarted := make(chan struct{}) + if !dispatcher.DispatchSized("alpha", 1, func() { + close(firstStarted) + <-releaseFirst + }) { + t.Fatal("dispatch first alpha item failed") + } + select { + case <-firstStarted: + case <-time.After(time.Second): + t.Fatal("timed out waiting for first alpha item") + } + if !dispatcher.DispatchSized("alpha", 1, func() {}) { + t.Fatal("dispatch second alpha item failed") + } + + alphaBlocked := make(chan bool, 1) + go func() { + alphaBlocked <- dispatcher.DispatchSized("alpha", 1, func() {}) + }() + select { + case <-alphaBlocked: + t.Fatal("alpha exceeded its per-source queue budget") + case <-time.After(100 * time.Millisecond): + } + + betaDone := make(chan bool, 1) + go func() { + betaDone <- dispatcher.DispatchSized("beta", 1, func() {}) + }() + select { + case ok := <-betaDone: + if !ok { + t.Fatal("beta dispatch failed") + } + case <-time.After(time.Second): + t.Fatal("a saturated alpha source starved the beta source") + } + + close(releaseFirst) + select { + case ok := <-alphaBlocked: + if !ok { + t.Fatal("blocked alpha dispatch failed after room became available") + } + case <-time.After(time.Second): + t.Fatal("blocked alpha dispatch was not released") + } +} diff --git a/logical_conn.go b/logical_conn.go index 8ec19d1..6a1f5c4 100644 --- a/logical_conn.go +++ b/logical_conn.go @@ -4,6 +4,7 @@ import ( "context" "errors" "net" + "sync" "sync/atomic" "time" ) @@ -18,6 +19,7 @@ type LogicalConn struct { transportState atomic.Pointer[clientConnTransportState] attachment atomic.Pointer[clientConnAttachmentState] inboundTransitionProfile atomic.Pointer[transportProtectionProfile] + transportLifecycleMu sync.Mutex } var errLogicalConnClientNil = errors.New("logical conn is nil") @@ -1118,15 +1120,23 @@ func (c *LogicalConn) transportConnSnapshotForInbound(conn net.Conn, remoteAddr } attached := false + var binding *transportBinding currentGeneration := c.transportGenerationSnapshot() if conn != nil { - binding := c.transportBindingSnapshot() - if binding != nil && binding.connSnapshot() == conn && c.transportAttachedSnapshot() && currentGeneration == generation { + currentBinding := c.transportBindingSnapshot() + if currentBinding != nil && currentBinding.connSnapshot() == conn && c.transportAttachedSnapshot() && currentGeneration == generation { + binding = currentBinding attached = true + } else { + // Keep stale inbound replies on the socket that delivered them. This + // binding intentionally has no queue/sender state and is never used + // by a logical send after IsCurrent rejects the old generation. + binding = newTransportBinding(conn, nil) } } else { current := c.CurrentTransportConn() if current != nil && currentGeneration == generation && transportConnAddrString(current.RemoteAddr()) == transportConnAddrString(remoteAddr) { + binding = current.binding attached = current.Attached() if !hasRuntimeConn { hasRuntimeConn = current.HasRuntimeConn() @@ -1138,6 +1148,7 @@ func (c *LogicalConn) transportConnSnapshotForInbound(conn net.Conn, remoteAddr logical: c, generation: generation, remoteAddr: remoteAddr, + binding: binding, attached: attached, hasRuntimeConn: hasRuntimeConn, } diff --git a/msg.go b/msg.go index 0171bcb..812b3a9 100644 --- a/msg.go +++ b/msg.go @@ -44,6 +44,7 @@ type Message struct { TransportConn *TransportConn ServerConn Client inboundTransportProfile *transportProtectionProfile + clientRoute clientSessionRoute TransferMsg Time time.Time inboundConn net.Conn @@ -73,6 +74,10 @@ type messageClientTransferSender interface { sendWithContextTimeout(context.Context, TransferMsg, time.Duration) (WaitMsg, error) } +type messageClientRouteTransferSender interface { + sendWithContextTimeoutAtRoute(context.Context, clientSessionRoute, TransferMsg, time.Duration) (WaitMsg, error) +} + type messageReplyWriteTimeoutProvider interface { ReplyWriteTimeout() time.Duration } @@ -140,6 +145,12 @@ func (m *Message) replyContext(ctx context.Context, value MsgVal) (err error) { if m.ServerConn == nil { return net.ErrClosed } + if m.clientRoute.bound() { + if sender, ok := m.ServerConn.(messageClientRouteTransferSender); ok { + _, err = sender.sendWithContextTimeoutAtRoute(ctx, m.clientRoute, reply, writeTimeout) + return err + } + } if sender, ok := m.ServerConn.(messageClientTransferSender); ok { _, err = sender.sendWithContextTimeout(ctx, reply, writeTimeout) } else { diff --git a/protocol_bounds_test.go b/protocol_bounds_test.go new file mode 100644 index 0000000..cb85a6b --- /dev/null +++ b/protocol_bounds_test.go @@ -0,0 +1,317 @@ +package notify + +import ( + "bytes" + "context" + "encoding/binary" + "errors" + "io" + "net" + "sync" + "testing" + "time" + + "b612.me/stario" +) + +type recordWriteCaptureStream struct { + mu sync.Mutex + buf bytes.Buffer + readDone chan struct{} + close sync.Once +} + +func newRecordWriteCaptureStream() *recordWriteCaptureStream { + return &recordWriteCaptureStream{ + readDone: make(chan struct{}), + } +} + +func (s *recordWriteCaptureStream) Read([]byte) (int, error) { + <-s.readDone + return 0, io.EOF +} + +func (s *recordWriteCaptureStream) Write(p []byte) (int, error) { + s.mu.Lock() + defer s.mu.Unlock() + return s.buf.Write(p) +} + +func (s *recordWriteCaptureStream) Close() error { + s.close.Do(func() { + close(s.readDone) + }) + return nil +} + +func (s *recordWriteCaptureStream) ID() string { return "record-capture" } +func (s *recordWriteCaptureStream) Channel() StreamChannel { return StreamRecordChannel } +func (s *recordWriteCaptureStream) Metadata() StreamMetadata { return nil } +func (s *recordWriteCaptureStream) Context() context.Context { return context.Background() } +func (s *recordWriteCaptureStream) LogicalConn() *LogicalConn { return nil } +func (s *recordWriteCaptureStream) TransportConn() *TransportConn { return nil } +func (s *recordWriteCaptureStream) TransportGeneration() uint64 { return 0 } +func (s *recordWriteCaptureStream) LocalAddr() net.Addr { return nil } +func (s *recordWriteCaptureStream) RemoteAddr() net.Addr { return nil } +func (s *recordWriteCaptureStream) CloseWrite() error { return nil } +func (s *recordWriteCaptureStream) Reset(error) error { return s.Close() } +func (s *recordWriteCaptureStream) SetDeadline(time.Time) error { return nil } +func (s *recordWriteCaptureStream) SetReadDeadline(time.Time) error { + return nil +} +func (s *recordWriteCaptureStream) SetWriteDeadline(time.Time) error { + return nil +} + +func (s *recordWriteCaptureStream) Bytes() []byte { + s.mu.Lock() + defer s.mu.Unlock() + return append([]byte(nil), s.buf.Bytes()...) +} + +func TestDedicatedRecordRejectsOversizedPayloadLength(t *testing.T) { + conn := &shortWriteBulkRecordConn{maxPerWrite: bulkDedicatedRecordMaxBytes + bulkDedicatedRecordHeaderLen + 1} + err := writeBulkDedicatedRecordWithDeadline(conn, make([]byte, bulkDedicatedRecordMaxBytes+1), time.Time{}) + if !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("writeBulkDedicatedRecordWithDeadline error = %v, want %v", err, errBulkFastPayloadInvalid) + } + if got := conn.buf.Len(); got != 0 { + t.Fatalf("oversized dedicated record wrote %d bytes, want 0", got) + } + + header := make([]byte, bulkDedicatedRecordHeaderLen) + copy(header[:4], bulkDedicatedRecordMagic) + binary.BigEndian.PutUint32(header[4:8], uint32(bulkDedicatedRecordMaxBytes+1)) + _, release, err := readBulkDedicatedRecordPooled(newBulkAttachScriptConn(header)) + if release != nil { + release() + t.Fatal("oversized dedicated record returned release callback") + } + if !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("readBulkDedicatedRecordPooled error = %v, want %v", err, errBulkFastPayloadInvalid) + } +} + +func TestDirectSignalFrameRejectsOversizedPayloadLength(t *testing.T) { + header := stario.NewQueue().BuildHeader(uint32(transportFrameMaxPayloadBytes + 1)) + _, err := readDirectSignalFramePayload(newBulkAttachScriptConn(header)) + if !errors.Is(err, stario.ErrQueueMessageTooLarge) { + t.Fatalf("readDirectSignalFramePayload error = %v, want %v", err, stario.ErrQueueMessageTooLarge) + } +} + +func TestTransferFrameRejectsOversizedPayloadLength(t *testing.T) { + stream := &transferWriteCountStream{} + var header [transferFrameHeaderSize]byte + binary.BigEndian.PutUint32(header[:], uint32(transferFrameMaxPayloadBytes+1)) + if _, err := stream.buf.Write(header[:]); err != nil { + t.Fatalf("seed transfer frame header failed: %v", err) + } + + _, err := readTransferFrame(stream) + if !errors.Is(err, errTransferFrameTooLarge) { + t.Fatalf("readTransferFrame error = %v, want %v", err, errTransferFrameTooLarge) + } +} + +func TestDedicatedBatchDecodersRejectOversizedWireCounts(t *testing.T) { + tooManyItems := make([]bulkDedicatedSendRequest, bulkDedicatedBatchMaxItems+1) + for i := range tooManyItems { + tooManyItems[i] = bulkDedicatedSendRequest{Type: bulkFastPayloadTypeData, Seq: uint64(i + 1)} + } + if _, err := encodeBulkDedicatedBatchPlain(1, tooManyItems); !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("encodeBulkDedicatedBatchPlain oversized item count error = %v, want %v", err, errBulkFastPayloadInvalid) + } + + tooManyGroups := make([]bulkDedicatedOutboundBatch, bulkDedicatedBatchMaxItems+1) + for i := range tooManyGroups { + tooManyGroups[i] = bulkDedicatedOutboundBatch{ + DataID: uint64(i + 1), + Items: []bulkDedicatedSendRequest{{ + Type: bulkFastPayloadTypeData, + Seq: 1, + }}, + } + } + if _, err := encodeBulkDedicatedBatchesPlain(tooManyGroups); !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("encodeBulkDedicatedBatchesPlain oversized group count error = %v, want %v", err, errBulkFastPayloadInvalid) + } + + batch := make([]byte, bulkDedicatedBatchHeaderLen) + copy(batch[:4], bulkDedicatedBatchMagic) + batch[4] = bulkDedicatedBatchVersion + binary.BigEndian.PutUint64(batch[8:16], 1) + binary.BigEndian.PutUint32(batch[16:20], uint32(bulkDedicatedBatchMaxItems+1)) + if _, _, matched, err := decodeBulkDedicatedBatchPlain(batch); !matched || !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("decodeBulkDedicatedBatchPlain matched=%v error=%v, want matched invalid", matched, err) + } + if err := walkDedicatedBulkInboundBatchPlain(batch, func(uint64, bulkDedicatedBatchItem) error { + t.Fatal("visit should not be called for oversized batch count") + return nil + }); !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("walkDedicatedBulkInboundBatchPlain error = %v, want %v", err, errBulkFastPayloadInvalid) + } + + superGroups := make([]byte, bulkDedicatedSuperBatchHeaderLen) + copy(superGroups[:4], bulkDedicatedSuperBatchMagic) + superGroups[4] = bulkDedicatedSuperBatchVersion + binary.BigEndian.PutUint32(superGroups[8:12], uint32(bulkDedicatedBatchMaxItems+1)) + if _, matched, err := decodeBulkDedicatedSuperBatchPlain(superGroups); !matched || !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("decodeBulkDedicatedSuperBatchPlain groups matched=%v error=%v, want matched invalid", matched, err) + } + if err := walkDedicatedBulkInboundSuperBatchPlain(superGroups, func(uint64, bulkDedicatedBatchItem) error { + t.Fatal("visit should not be called for oversized super-batch group count") + return nil + }); !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("walkDedicatedBulkInboundSuperBatchPlain groups error = %v, want %v", err, errBulkFastPayloadInvalid) + } + + superItems := make([]byte, bulkDedicatedSuperBatchHeaderLen+bulkDedicatedSuperBatchGroupHeaderLen) + copy(superItems[:4], bulkDedicatedSuperBatchMagic) + superItems[4] = bulkDedicatedSuperBatchVersion + binary.BigEndian.PutUint32(superItems[8:12], 1) + binary.BigEndian.PutUint64(superItems[12:20], 1) + binary.BigEndian.PutUint32(superItems[20:24], uint32(bulkDedicatedBatchMaxItems+1)) + if _, matched, err := decodeBulkDedicatedSuperBatchPlain(superItems); !matched || !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("decodeBulkDedicatedSuperBatchPlain items matched=%v error=%v, want matched invalid", matched, err) + } + if err := walkDedicatedBulkInboundSuperBatchPlain(superItems, func(uint64, bulkDedicatedBatchItem) error { + t.Fatal("visit should not be called for oversized super-batch item count") + return nil + }); !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("walkDedicatedBulkInboundSuperBatchPlain items error = %v, want %v", err, errBulkFastPayloadInvalid) + } +} + +func TestSharedFastBatchDecodersRejectOversizedWireCounts(t *testing.T) { + tooManyBulkFrames := make([]bulkFastFrame, bulkFastBatchMaxItems+1) + for i := range tooManyBulkFrames { + tooManyBulkFrames[i] = bulkFastFrame{Type: bulkFastPayloadTypeData, DataID: 1, Seq: uint64(i + 1)} + } + if _, err := encodeBulkFastBatchPlain(tooManyBulkFrames); !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("encodeBulkFastBatchPlain oversized count error = %v, want %v", err, errBulkFastPayloadInvalid) + } + + bulkBatch := make([]byte, bulkFastBatchHeaderLen) + copy(bulkBatch[:4], bulkFastBatchMagic) + bulkBatch[4] = bulkFastBatchVersion + binary.BigEndian.PutUint32(bulkBatch[8:12], uint32(bulkFastBatchMaxItems+1)) + if matched, err := walkBulkFastBatchPlain(bulkBatch, func(bulkFastFrame) error { + t.Fatal("bulk batch visitor should not be called for oversized count") + return nil + }); !matched || !errors.Is(err, errBulkFastPayloadInvalid) { + t.Fatalf("walkBulkFastBatchPlain matched=%v error=%v, want matched invalid", matched, err) + } + + tooManyStreamFrames := make([]streamFastDataFrame, streamFastBatchMaxItems+1) + for i := range tooManyStreamFrames { + tooManyStreamFrames[i] = streamFastDataFrame{DataID: 1, Seq: uint64(i + 1)} + } + if _, err := encodeStreamFastBatchPlain(tooManyStreamFrames); !errors.Is(err, errStreamFastPayloadInvalid) { + t.Fatalf("encodeStreamFastBatchPlain oversized count error = %v, want %v", err, errStreamFastPayloadInvalid) + } + + streamBatch := make([]byte, streamFastBatchHeaderLen) + copy(streamBatch[:4], streamFastBatchMagic) + streamBatch[4] = streamFastBatchVersion + binary.BigEndian.PutUint32(streamBatch[8:12], uint32(streamFastBatchMaxItems+1)) + if matched, err := walkStreamFastBatchPlain(streamBatch, func(streamFastDataFrame) error { + t.Fatal("stream batch visitor should not be called for oversized count") + return nil + }); !matched || !errors.Is(err, errStreamFastPayloadInvalid) { + t.Fatalf("walkStreamFastBatchPlain matched=%v error=%v, want matched invalid", matched, err) + } +} + +func TestRecordStreamRejectsPayloadLargerThanUnackedWindow(t *testing.T) { + record := &recordStream{ + cfg: recordConfig{ + MaxUnackedBytes: 4, + }, + } + _, err := record.WriteRecord(context.Background(), []byte("12345")) + if !errors.Is(err, errRecordPayloadTooLarge) { + t.Fatalf("WriteRecord error = %v, want %v", err, errRecordPayloadTooLarge) + } +} + +func TestRecordOptionsCapBatchCountsToWireLimit(t *testing.T) { + opt := normalizeRecordOpenOptions(RecordOpenOptions{ + MaxBatchRecords: recordMaxBatchRecords + 100, + MaxBatchBytes: transferFrameMaxPayloadBytes * 2, + MaxUnackedBytes: transferFrameMaxPayloadBytes * 2, + }) + if got, want := opt.MaxBatchRecords, recordMaxBatchRecords; got != want { + t.Fatalf("MaxBatchRecords = %d, want %d", got, want) + } + if got, want := opt.MaxBatchBytes, recordMaxBatchPayloadBytes; got != want { + t.Fatalf("MaxBatchBytes = %d, want %d", got, want) + } + if got, want := opt.MaxUnackedBytes, transferFrameMaxPayloadBytes*2; got != want { + t.Fatalf("MaxUnackedBytes = %d, want %d", got, want) + } +} + +func TestRecordWriterFlushesBeforeAppendingPastBatchByteLimit(t *testing.T) { + stream := newRecordWriteCaptureStream() + record, err := WrapStreamAsRecord(stream, RecordOpenOptions{ + MaxBatchRecords: defaultRecordMaxBatchRecords, + MaxBatchBytes: 10, + MaxBatchDelay: time.Hour, + MaxUnackedRecords: 16, + MaxUnackedBytes: 1024, + }) + if err != nil { + t.Fatalf("WrapStreamAsRecord failed: %v", err) + } + defer record.Close() + + if _, err := record.WriteRecord(context.Background(), []byte("123456")); err != nil { + t.Fatalf("first WriteRecord failed: %v", err) + } + if _, err := record.WriteRecord(context.Background(), []byte("abcdef")); err != nil { + t.Fatalf("second WriteRecord failed: %v", err) + } + flushCtx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := record.Flush(flushCtx); err != nil { + t.Fatalf("Flush failed: %v", err) + } + + if got, want := countTransferFrames(stream.Bytes()), 2; got != want { + t.Fatalf("transfer frame count = %d, want %d", got, want) + } +} + +func TestRecordBatchDecodersRejectOversizedWireCountsAndLengths(t *testing.T) { + v1 := makeRecordBatchFrameHeaderForBoundsTest(recordFrameVersionV1, defaultRecordMaxBatchRecords+1) + if _, err := decodeRecordFrame(v1); !errors.Is(err, errRecordFrameInvalid) { + t.Fatalf("decodeRecordFrame oversized v1 count error = %v, want %v", err, errRecordFrameInvalid) + } + + v2 := makeRecordBatchFrameHeaderForBoundsTest(recordFrameVersionV2, defaultRecordMaxBatchRecords+1) + if _, err := decodeRecordFrame(v2); !errors.Is(err, errRecordFrameInvalid) { + t.Fatalf("decodeRecordFrame oversized v2 count error = %v, want %v", err, errRecordFrameInvalid) + } + + v2LongItem := makeRecordBatchFrameHeaderForBoundsTest(recordFrameVersionV2, 1) + v2LongItem = append(v2LongItem, 0xff, 0xff, 0xff, 0xff) + if _, err := decodeRecordFrame(v2LongItem); !errors.Is(err, errRecordFrameInvalid) { + t.Fatalf("decodeRecordFrame oversized v2 item length error = %v, want %v", err, errRecordFrameInvalid) + } +} + +func makeRecordBatchFrameHeaderForBoundsTest(version uint8, count int) []byte { + headerSize := recordBatchHeaderV1Size + if version == recordFrameVersionV2 { + headerSize = recordBatchHeaderV2Size + } + frame := make([]byte, recordFrameHeaderSize+headerSize) + copy(frame[:4], recordFrameMagic) + frame[4] = version + frame[5] = recordFrameTypeBatch + binary.BigEndian.PutUint16(frame[8:10], uint16(count)) + binary.BigEndian.PutUint64(frame[10:18], 1) + return frame +} diff --git a/record_benchmark_test.go b/record_benchmark_test.go new file mode 100644 index 0000000..47bbc1c --- /dev/null +++ b/record_benchmark_test.go @@ -0,0 +1,40 @@ +package notify + +import ( + "context" + "testing" +) + +type discardRecordWriteStream struct{ *recordWriteCaptureStream } + +func (s *discardRecordWriteStream) Write(p []byte) (int, error) { return len(p), nil } + +func BenchmarkRecordWriteFlush64(b *testing.B) { + s := &discardRecordWriteStream{newRecordWriteCaptureStream()} + rs, err := WrapStreamAsRecord(s, RecordOpenOptions{}) + if err != nil { + b.Fatal(err) + } + r := rs.(*recordStream) + b.Cleanup(func() { _ = r.Close() }) + payload := make([]byte, 1024) + ctx := context.Background() + b.SetBytes(64 * int64(len(payload))) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + var last uint64 + for j := 0; j < 64; j++ { + last, err = r.WriteRecord(ctx, payload) + if err != nil { + b.Fatal(err) + } + } + if err := r.Flush(ctx); err != nil { + b.Fatal(err) + } + if err := r.handleAckFrame(last); err != nil { + b.Fatal(err) + } + } +} diff --git a/record_codec.go b/record_codec.go index 9c6bc19..2d6ce3b 100644 --- a/record_codec.go +++ b/record_codec.go @@ -12,6 +12,7 @@ const ( recordFrameTypeBatch uint8 = 1 recordFrameTypeAck uint8 = 2 recordFrameTypeError uint8 = 3 + recordFrameTypeFIN uint8 = 4 recordFrameHeaderSize = 8 recordBatchHeaderV1Size = 10 recordBatchHeaderV2Size = 18 @@ -33,6 +34,7 @@ type recordFrame struct { Type uint8 Batch []recordOutboundMessage AckSeq uint64 + FinalSeq uint64 Failure RecordFailure Retryable bool } @@ -41,8 +43,11 @@ func encodeRecordBatchFrame(batch []recordOutboundMessage, ackSeq uint64, useV2 if len(batch) == 0 { return nil, nil } + if len(batch) > recordMaxBatchRecords { + return nil, errRecordFrameInvalid + } firstSeq := batch[0].Seq - if firstSeq == 0 { + if firstSeq == 0 || uint64(len(batch)-1) > ^uint64(0)-firstSeq { return nil, errRecordSeqInvalid } version := uint8(recordFrameVersionV1) @@ -57,6 +62,9 @@ func encodeRecordBatchFrame(batch []recordOutboundMessage, ackSeq uint64, useV2 if item.Seq != wantSeq { return nil, errRecordSeqInvalid } + if len(item.Payload) > recordMaxPayloadBytes || size > transferFrameMaxPayloadBytes-4-len(item.Payload) { + return nil, errRecordFrameInvalid + } size += 4 + len(item.Payload) } frame := make([]byte, size) @@ -87,13 +95,21 @@ func encodeRecordAckFrame(ackSeq uint64) ([]byte, error) { return frame, nil } +func encodeRecordFINFrame(finalSeq uint64) []byte { + frame, _ := encodeRecordAckFrame(finalSeq) + frame[5] = recordFrameTypeFIN + return frame +} + func encodeRecordErrorFrame(failure RecordFailure) ([]byte, error) { if failure.FailedSeq == 0 { return nil, errRecordSeqInvalid } - codeBytes := []byte(failure.Code) - msgBytes := []byte(failure.Message) - frame := make([]byte, recordFrameHeaderSize+recordErrorHeaderSize+len(codeBytes)+len(msgBytes)) + const headerSize = recordFrameHeaderSize + recordErrorHeaderSize + if len(failure.Code) > int(^uint16(0)) || len(failure.Message) > transferFrameMaxPayloadBytes-headerSize-len(failure.Code) { + return nil, errRecordFrameInvalid + } + frame := make([]byte, headerSize+len(failure.Code)+len(failure.Message)) copy(frame[:4], recordFrameMagic) frame[4] = recordFrameVersionV1 frame[5] = recordFrameTypeError @@ -101,12 +117,12 @@ func encodeRecordErrorFrame(failure RecordFailure) ([]byte, error) { frame[6] = 1 } binary.BigEndian.PutUint64(frame[8:16], failure.FailedSeq) - binary.BigEndian.PutUint16(frame[16:18], uint16(len(codeBytes))) - binary.BigEndian.PutUint32(frame[18:22], uint32(len(msgBytes))) + binary.BigEndian.PutUint16(frame[16:18], uint16(len(failure.Code))) + binary.BigEndian.PutUint32(frame[18:22], uint32(len(failure.Message))) offset := recordFrameHeaderSize + recordErrorHeaderSize - copy(frame[offset:offset+len(codeBytes)], codeBytes) - offset += len(codeBytes) - copy(frame[offset:offset+len(msgBytes)], msgBytes) + copy(frame[offset:], failure.Code) + offset += len(failure.Code) + copy(frame[offset:], failure.Message) return frame, nil } @@ -116,6 +132,12 @@ func decodeRecordFrame(payload []byte) (recordFrame, error) { } version := payload[4] frameType := payload[5] + if frameType == recordFrameTypeFIN && version == recordFrameVersionV1 { + if len(payload) != recordFrameHeaderSize+8 { + return recordFrame{}, errRecordFrameInvalid + } + return recordFrame{Version: version, Type: frameType, FinalSeq: binary.BigEndian.Uint64(payload[8:16])}, nil + } switch version { case recordFrameVersionV1: switch frameType { @@ -174,20 +196,24 @@ func decodeRecordBatchFrameV1(payload []byte) (recordFrame, error) { } count := int(binary.BigEndian.Uint16(payload[8:10])) firstSeq := binary.BigEndian.Uint64(payload[10:18]) - if count <= 0 || firstSeq == 0 { + if count <= 0 || firstSeq == 0 || uint64(count-1) > ^uint64(0)-firstSeq { return recordFrame{}, errRecordFrameInvalid } offset := recordFrameHeaderSize + recordBatchHeaderV1Size + if count > (len(payload)-offset)/4 { + return recordFrame{}, errRecordFrameInvalid + } batch := make([]recordOutboundMessage, 0, count) for index := 0; index < count; index++ { - if offset+4 > len(payload) { + if len(payload)-offset < 4 { return recordFrame{}, errRecordFrameInvalid } - itemLen := int(binary.BigEndian.Uint32(payload[offset : offset+4])) + wireItemLen := binary.BigEndian.Uint32(payload[offset : offset+4]) offset += 4 - if itemLen < 0 || offset+itemLen > len(payload) { + if uint64(wireItemLen) > uint64(len(payload)-offset) { return recordFrame{}, errRecordFrameInvalid } + itemLen := int(wireItemLen) item := recordOutboundMessage{ Seq: firstSeq + uint64(index), Payload: append([]byte(nil), payload[offset:offset+itemLen]...), @@ -212,20 +238,24 @@ func decodeRecordBatchFrameV2(payload []byte) (recordFrame, error) { count := int(binary.BigEndian.Uint16(payload[8:10])) firstSeq := binary.BigEndian.Uint64(payload[10:18]) ackSeq := binary.BigEndian.Uint64(payload[18:26]) - if count <= 0 || firstSeq == 0 { + if count <= 0 || firstSeq == 0 || uint64(count-1) > ^uint64(0)-firstSeq { return recordFrame{}, errRecordFrameInvalid } offset := recordFrameHeaderSize + recordBatchHeaderV2Size + if count > (len(payload)-offset)/4 { + return recordFrame{}, errRecordFrameInvalid + } batch := make([]recordOutboundMessage, 0, count) for index := 0; index < count; index++ { - if offset+4 > len(payload) { + if len(payload)-offset < 4 { return recordFrame{}, errRecordFrameInvalid } - itemLen := int(binary.BigEndian.Uint32(payload[offset : offset+4])) + wireItemLen := binary.BigEndian.Uint32(payload[offset : offset+4]) offset += 4 - if itemLen < 0 || offset+itemLen > len(payload) { + if uint64(wireItemLen) > uint64(len(payload)-offset) { return recordFrame{}, errRecordFrameInvalid } + itemLen := int(wireItemLen) item := recordOutboundMessage{ Seq: firstSeq + uint64(index), Payload: append([]byte(nil), payload[offset:offset+itemLen]...), @@ -250,9 +280,12 @@ func decodeRecordErrorFrame(payload []byte) (recordFrame, error) { } failedSeq := binary.BigEndian.Uint64(payload[8:16]) codeLen := int(binary.BigEndian.Uint16(payload[16:18])) - msgLen := int(binary.BigEndian.Uint32(payload[18:22])) + wireMsgLen := binary.BigEndian.Uint32(payload[18:22]) offset := recordFrameHeaderSize + recordErrorHeaderSize - if failedSeq == 0 || offset+codeLen+msgLen != len(payload) { + if failedSeq == 0 || len(payload)-offset < codeLen { + return recordFrame{}, errRecordFrameInvalid + } + if uint64(len(payload)-offset-codeLen) != uint64(wireMsgLen) { return recordFrame{}, errRecordFrameInvalid } failure := RecordFailure{ diff --git a/record_lifecycle.go b/record_lifecycle.go new file mode 100644 index 0000000..d53e746 --- /dev/null +++ b/record_lifecycle.go @@ -0,0 +1,184 @@ +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 +} diff --git a/record_negotiation.go b/record_negotiation.go index 891f234..ccbf6b2 100644 --- a/record_negotiation.go +++ b/record_negotiation.go @@ -1,9 +1,11 @@ package notify const ( - recordStreamMetadataCapBatchAckKey = "_notify.record_cap_batch_ack" - recordStreamMetadataUseBatchAckKey = "_notify.record_use_batch_ack" - recordStreamMetadataEnabledValue = "1" + recordStreamMetadataCapBatchAckKey = "_notify.record_cap_batch_ack" + recordStreamMetadataUseBatchAckKey = "_notify.record_use_batch_ack" + recordStreamMetadataCapHalfCloseKey = "_notify.record_cap_half_close" + recordStreamMetadataUseHalfCloseKey = "_notify.record_use_half_close" + recordStreamMetadataEnabledValue = "1" ) func advertiseRecordStreamOpenMetadata(metadata StreamMetadata) StreamMetadata { @@ -11,7 +13,9 @@ func advertiseRecordStreamOpenMetadata(metadata StreamMetadata) StreamMetadata { if metadata == nil { metadata = make(StreamMetadata, 1) } + delete(metadata, recordStreamMetadataUseHalfCloseKey) metadata[recordStreamMetadataCapBatchAckKey] = recordStreamMetadataEnabledValue + metadata[recordStreamMetadataCapHalfCloseKey] = recordStreamMetadataEnabledValue return metadata } @@ -20,13 +24,17 @@ func negotiateRecordStreamOpenMetadata(channel StreamChannel, metadata StreamMet if normalizeStreamChannel(channel) != StreamRecordChannel { return metadata, nil } - if metadata[recordStreamMetadataCapBatchAckKey] != recordStreamMetadataEnabledValue { - return metadata, nil + response := make(StreamMetadata) + if metadata[recordStreamMetadataCapBatchAckKey] == recordStreamMetadataEnabledValue { + metadata[recordStreamMetadataUseBatchAckKey] = recordStreamMetadataEnabledValue + response[recordStreamMetadataUseBatchAckKey] = recordStreamMetadataEnabledValue } - metadata[recordStreamMetadataUseBatchAckKey] = recordStreamMetadataEnabledValue - return metadata, StreamMetadata{ - recordStreamMetadataUseBatchAckKey: recordStreamMetadataEnabledValue, + delete(metadata, recordStreamMetadataUseHalfCloseKey) + if metadata[recordStreamMetadataCapHalfCloseKey] == recordStreamMetadataEnabledValue { + metadata[recordStreamMetadataUseHalfCloseKey] = recordStreamMetadataEnabledValue + response[recordStreamMetadataUseHalfCloseKey] = recordStreamMetadataEnabledValue } + return metadata, response } func mergeStreamMetadata(base StreamMetadata, overlay StreamMetadata) StreamMetadata { @@ -46,3 +54,7 @@ func mergeStreamMetadata(base StreamMetadata, overlay StreamMetadata) StreamMeta func recordStreamUseBatchAck(metadata StreamMetadata) bool { return metadata[recordStreamMetadataUseBatchAckKey] == recordStreamMetadataEnabledValue } + +func recordStreamUseHalfClose(metadata StreamMetadata) bool { + return metadata[recordStreamMetadataUseHalfCloseKey] == recordStreamMetadataEnabledValue +} diff --git a/record_network_test.go b/record_network_test.go new file mode 100644 index 0000000..1378887 --- /dev/null +++ b/record_network_test.go @@ -0,0 +1,177 @@ +package notify + +import ( + "bytes" + "context" + "encoding/binary" + "fmt" + "io" + "net" + "sync" + "testing" + "time" +) + +type delayedRecordPacket struct { + data []byte + ready time.Time +} + +// Delay packets in a bounded pipeline so latency does not serialize every write. +func startRecordLinkProxy(t *testing.T, upstream string, mbps int, rtt time.Duration) string { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + var workers sync.WaitGroup + workers.Add(1) + go func() { + defer workers.Done() + left, err := listener.Accept() + if err != nil { + return + } + defer left.Close() + right, err := (&net.Dialer{}).DialContext(ctx, "tcp", upstream) + if err != nil { + return + } + defer right.Close() + stopConn := context.AfterFunc(ctx, func() { _ = left.Close(); _ = right.Close() }) + defer stopConn() + var relays sync.WaitGroup + relay := func(dst, src net.Conn) { + defer relays.Done() + packets := make(chan delayedRecordPacket, 128) + var reader sync.WaitGroup + reader.Add(1) + go func() { + defer reader.Done() + defer close(packets) + buf := make([]byte, 16*1024) + for { + n, err := src.Read(buf) + if n > 0 { + packet := delayedRecordPacket{append([]byte(nil), buf[:n]...), time.Now().Add(rtt / 2)} + select { + case packets <- packet: + case <-ctx.Done(): + return + } + } + if err != nil { + return + } + } + }() + defer reader.Wait() + defer cancel() + var next time.Time + for packet := range packets { + if next.Before(time.Now()) { + next = time.Now() + } + next = next.Add(time.Duration(len(packet.data)) * time.Second / time.Duration(mbps*1000*1000/8)) + ready := packet.ready + if next.After(ready) { + ready = next + } + timer := time.NewTimer(time.Until(ready)) + select { + case <-ctx.Done(): + timer.Stop() + return + case <-timer.C: + } + if err := writeFullToConn(dst, packet.data); err != nil { + return + } + } + } + relays.Add(2) + go relay(right, left) + go relay(left, right) + relays.Wait() + }() + t.Cleanup(func() { cancel(); _ = listener.Close(); workers.Wait() }) + return listener.Addr().String() +} + +func TestRecordTCPDelayedBandwidth(t *testing.T) { + for _, tc := range []struct { + mbps int + rtt time.Duration + }{{10, 80 * time.Millisecond}, {50, 160 * time.Millisecond}, {100, 80 * time.Millisecond}} { + t.Run(fmt.Sprintf("%dMbps-%s", tc.mbps, tc.rtt), func(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatal(err) + } + const count, size = 512, 2048 + handlerDone := make(chan error, 1) + server.SetRecordStreamHandler(func(info RecordAcceptInfo) error { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + for i := 0; i < count; i++ { + msg, err := info.RecordStream.ReadRecord(ctx) + if err != nil { + handlerDone <- err + return err + } + want := bytes.Repeat([]byte{byte(i)}, size) + binary.BigEndian.PutUint64(want[:8], uint64(i)) + if msg.Seq != uint64(i+1) || !bytes.Equal(msg.Payload, want) { + err = fmt.Errorf("record %d corrupted", i) + handlerDone <- err + return err + } + if err := info.RecordStream.AckRecord(msg.Seq); err != nil { + handlerDone <- err + return err + } + } + handlerDone <- nil + return nil + }) + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = server.Stop() }) + proxy := startRecordLinkProxy(t, server.listener.Addr().String(), tc.mbps, tc.rtt) + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatal(err) + } + if err := client.Connect("tcp", proxy); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Stop() }) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + r, err := client.OpenRecordStream(ctx, RecordOpenOptions{}) + if err != nil { + t.Fatal(err) + } + started := time.Now() + for i := 0; i < count; i++ { + payload := bytes.Repeat([]byte{byte(i)}, size) + binary.BigEndian.PutUint64(payload[:8], uint64(i)) + if _, err := r.WriteRecord(ctx, payload); err != nil { + t.Fatal(err) + } + } + if acked, err := r.Barrier(ctx); err != nil || acked != count { + t.Fatalf("barrier ack=%d err=%v", acked, err) + } + if err := <-handlerDone; err != nil { + t.Fatal(err) + } + if err := r.Close(); err != nil && err != io.EOF { + t.Fatal(err) + } + t.Logf("verified %d records, %d bytes in %s", count, count*size, time.Since(started)) + }) + } +} diff --git a/record_protocol_lifecycle_test.go b/record_protocol_lifecycle_test.go new file mode 100644 index 0000000..0aeaee5 --- /dev/null +++ b/record_protocol_lifecycle_test.go @@ -0,0 +1,318 @@ +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) + } +} diff --git a/record_regression_test.go b/record_regression_test.go new file mode 100644 index 0000000..25a60ea --- /dev/null +++ b/record_regression_test.go @@ -0,0 +1,373 @@ +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) + } +} diff --git a/record_reset.go b/record_reset.go new file mode 100644 index 0000000..3ce03ae --- /dev/null +++ b/record_reset.go @@ -0,0 +1,23 @@ +package notify + +import "errors" + +// Reset control bypasses record receive backpressure. Keep the typed cause in +// that message as well as the in-band error frame used by older peers. +func (s *streamHandle) recordResetFailure() *RecordFailure { + if s.Channel() != StreamRecordChannel { + return nil + } + var failure RecordFailure + if errors.As(s.resetErrSnapshot(), &failure) { + return &failure + } + return nil +} + +func (r StreamResetRequest) resetError(channel StreamChannel) error { + if channel == StreamRecordChannel && r.RecordFailure != nil { + return *r.RecordFailure + } + return streamRemoteResetError(r.Error) +} diff --git a/record_stream.go b/record_stream.go index ffa9d04..f193bad 100644 --- a/record_stream.go +++ b/record_stream.go @@ -27,6 +27,9 @@ const ( defaultRecordInboundQueueLimit = 128 defaultRecordAckEveryRecords = 64 defaultRecordAckDelay = time.Millisecond + recordMaxBatchRecords = 1<<16 - 1 + recordMaxPayloadBytes = transferFrameMaxPayloadBytes - recordFrameHeaderSize - recordBatchHeaderV2Size - 4 + recordMaxBatchPayloadBytes = transferFrameMaxPayloadBytes - recordFrameHeaderSize - recordBatchHeaderV2Size - 4*recordMaxBatchRecords ) type RecordFailure struct { @@ -100,11 +103,14 @@ type recordConfig struct { InboundQueueLimit int AckEveryRecords int AckDelay time.Duration + CloseTimeout time.Duration } type recordFlushRequest struct { + ctx context.Context targetSeq uint64 forceAck bool + closeMode recordCloseMode done chan error } @@ -123,18 +129,26 @@ type recordObservability struct { } type recordStream struct { - stream Stream - ctx context.Context - cancel context.CancelFunc - cfg recordConfig - writeMu sync.Mutex - sendCh chan recordOutboundMessage - flushCh chan recordFlushRequest - recvCh chan RecordMessage - ackCh chan struct{} - readerCh chan struct{} - useBatchAck bool - obs recordObservability + stream Stream + ctx context.Context + cancel context.CancelFunc + cfg recordConfig + writeMu sync.Mutex + sendCh chan recordOutboundMessage + sendReady chan struct{} + flushCh chan recordFlushRequest + recvCh chan RecordMessage + ackCh chan struct{} + readerCh chan struct{} + writerCh chan struct{} + useBatchAck bool + useHalfClose bool + recvCloseOnce sync.Once + obs recordObservability + closeMu sync.Mutex + closeDone bool + closeErr error + halfClosed bool mu sync.Mutex @@ -160,9 +174,10 @@ type recordStream struct { inboundAckSentSeq uint64 maxPendingApply int - remoteClosed bool - readErr error - terminalErr error + remoteClosed bool + inboundClosed bool + readErr error + terminalErr error } var ( @@ -170,6 +185,7 @@ var ( errRecordRuntimeNil = errors.New("record runtime is nil") errRecordHandlerNotConfigured = errors.New("record handler is not configured") errRecordWriteClosed = errors.New("record stream write side is closed") + errRecordPayloadTooLarge = errors.New("record payload too large") errRecordSeqNotReceived = errors.New("record sequence not received") ) @@ -177,9 +193,15 @@ func normalizeRecordOpenOptions(opt RecordOpenOptions) RecordOpenOptions { if opt.MaxBatchRecords <= 0 { opt.MaxBatchRecords = defaultRecordMaxBatchRecords } + if opt.MaxBatchRecords > recordMaxBatchRecords { + opt.MaxBatchRecords = recordMaxBatchRecords + } if opt.MaxBatchBytes <= 0 { opt.MaxBatchBytes = defaultRecordMaxBatchBytes } + if opt.MaxBatchBytes > recordMaxBatchPayloadBytes { + opt.MaxBatchBytes = recordMaxBatchPayloadBytes + } if opt.MaxBatchDelay <= 0 { opt.MaxBatchDelay = defaultRecordMaxBatchDelay } @@ -212,6 +234,7 @@ func recordConfigFromOptions(opt RecordOpenOptions) recordConfig { InboundQueueLimit: opt.InboundQueueLimit, AckEveryRecords: opt.AckEveryRecords, AckDelay: opt.AckDelay, + CloseTimeout: opt.Stream.WriteTimeout, } } @@ -232,16 +255,19 @@ func WrapStreamAsRecord(stream Stream, opt RecordOpenOptions) (RecordStream, err } ctx, cancel := context.WithCancel(parent) record := &recordStream{ - stream: stream, - ctx: ctx, - cancel: cancel, - cfg: recordConfigFromOptions(opt), - sendCh: make(chan recordOutboundMessage, opt.MaxBatchRecords*2), - flushCh: make(chan recordFlushRequest), - recvCh: make(chan RecordMessage, opt.InboundQueueLimit), - ackCh: make(chan struct{}, 1), - readerCh: make(chan struct{}), - useBatchAck: recordStreamUseBatchAck(stream.Metadata()), + stream: stream, + ctx: ctx, + cancel: cancel, + cfg: recordConfigFromOptions(opt), + sendCh: make(chan recordOutboundMessage, opt.MaxBatchRecords*2), + sendReady: make(chan struct{}, 1), + flushCh: make(chan recordFlushRequest), + recvCh: make(chan RecordMessage, opt.InboundQueueLimit), + ackCh: make(chan struct{}, 1), + readerCh: make(chan struct{}), + writerCh: make(chan struct{}), + useBatchAck: recordStreamUseBatchAck(stream.Metadata()), + useHalfClose: recordStreamUseHalfClose(stream.Metadata()), stateNotify: make(chan struct{}), outstandingSizes: make(map[uint64]int), @@ -283,6 +309,10 @@ func (r *recordStream) WriteRecord(ctx context.Context, payload []byte) (uint64, size := len(payload) for { r.mu.Lock() + if err := ctx.Err(); err != nil { + r.mu.Unlock() + return 0, err + } if err := r.streamErrorLocked(); err != nil { r.mu.Unlock() return 0, err @@ -291,7 +321,15 @@ func (r *recordStream) WriteRecord(ctx context.Context, payload []byte) (uint64, r.mu.Unlock() return 0, errRecordWriteClosed } - if r.outstandingRecords >= r.cfg.MaxUnackedRecords || r.outstandingBytes+size > r.cfg.MaxUnackedBytes { + if size > recordMaxPayloadBytes { + r.mu.Unlock() + return 0, fmt.Errorf("%w: size=%d max_record_payload=%d", errRecordPayloadTooLarge, size, recordMaxPayloadBytes) + } + if size > r.cfg.MaxUnackedBytes { + r.mu.Unlock() + return 0, fmt.Errorf("%w: size=%d max_unacked_bytes=%d", errRecordPayloadTooLarge, size, r.cfg.MaxUnackedBytes) + } + if r.outstandingRecords >= r.cfg.MaxUnackedRecords || size > r.cfg.MaxUnackedBytes-r.outstandingBytes || len(r.sendCh) == cap(r.sendCh) { wait := r.stateNotify r.mu.Unlock() select { @@ -300,9 +338,14 @@ func (r *recordStream) WriteRecord(ctx context.Context, payload []byte) (uint64, case <-ctx.Done(): return 0, ctx.Err() case <-wait: + case <-r.sendReady: } continue } + if r.nextOutboundSeq == ^uint64(0) { + r.mu.Unlock() + return 0, errRecordSeqInvalid + } r.nextOutboundSeq++ msg := recordOutboundMessage{ Seq: r.nextOutboundSeq, @@ -311,22 +354,12 @@ func (r *recordStream) WriteRecord(ctx context.Context, payload []byte) (uint64, r.outstandingRecords++ r.outstandingBytes += size r.outstandingSizes[msg.Seq] = size - select { - case <-r.ctx.Done(): - r.rollbackReservedOutboundLocked(msg.Seq) - err := r.streamErrorLocked() - r.mu.Unlock() - return 0, err - case <-ctx.Done(): - r.rollbackReservedOutboundLocked(msg.Seq) - r.mu.Unlock() - return 0, ctx.Err() - case r.sendCh <- msg: - r.enqueuedOutboundSeq = msg.Seq - r.signalStateLocked() - r.mu.Unlock() - return msg.Seq, nil - } + // All producers hold mu; only the consumer can change the checked capacity. + r.sendCh <- msg + r.enqueuedOutboundSeq = msg.Seq + r.signalStateLocked() + r.mu.Unlock() + return msg.Seq, nil } } @@ -341,6 +374,7 @@ func (r *recordStream) Flush(ctx context.Context) error { return err } req := recordFlushRequest{ + ctx: ctx, targetSeq: r.flushTargetSeq(), done: make(chan error, 1), } @@ -474,44 +508,34 @@ func (r *recordStream) FailRecord(seq uint64, failure RecordFailure) error { if failure.Code == "" { failure.Code = RecordErrorCodeApplyFailed } - err := r.sendFailureFrame(failure) - if err != nil { - return err - } r.setTerminalError(failure) - return r.stream.Reset(failure) + return r.notifyFailureAndAbort(failure) } -func (r *recordStream) CloseWrite() error { - if r == nil { - return errRecordStreamNil +func (r *recordStream) notifyFailureAndAbort(failure RecordFailure) error { + ctx, cancel := context.WithTimeout(context.Background(), r.abortTimeout()) + defer cancel() + deadline, _ := ctx.Deadline() + _ = r.stream.SetWriteDeadline(deadline) + done := make(chan error, 1) + go func() { done <- r.sendFailureFrame(failure) }() + var sendErr error + select { + case sendErr = <-done: + case <-ctx.Done(): + sendErr = ctx.Err() } - if err := r.Flush(context.Background()); err != nil { - return err - } - if err := r.flushAckNow(); err != nil { - return err - } - r.mu.Lock() - r.outboundClosed = true - r.signalStateLocked() - r.mu.Unlock() - return r.stream.CloseWrite() -} - -func (r *recordStream) Close() error { - if r == nil { - return nil - } - _ = r.flushAckNow() - r.cancel() - return r.stream.Close() + r.abortUnderlyingStream(failure) + return sendErr } func (r *recordStream) Reset(err error) error { if r == nil { return nil } + if err == nil { + err = io.ErrClosedPipe + } r.setTerminalError(err) return r.stream.Reset(err) } @@ -553,6 +577,7 @@ func (r *recordStream) waitAckedAtLeast(ctx context.Context, target uint64) erro } func (r *recordStream) writerLoop() { + defer close(r.writerCh) var ( batch []recordOutboundMessage batches int @@ -586,6 +611,8 @@ func (r *recordStream) writerLoop() { } ackTimerCh = nil } + defer stopBatchTimer() + defer stopAckTimer() scheduleAck := func(hasPendingBatch bool, force bool) (uint64, bool) { ackSeq := r.pendingAckSeq() if ackSeq == 0 { @@ -654,6 +681,23 @@ func (r *recordStream) writerLoop() { } return nil } + appendBatch := func(req recordOutboundMessage) (bool, error) { + if len(batch) > 0 && (batches >= r.cfg.MaxBatchRecords || bytes+len(req.Payload) > r.cfg.MaxBatchBytes) { + if err := flushBatch(); err != nil { + return false, err + } + } + batch = append(batch, req) + batches++ + bytes += len(req.Payload) + if batches >= r.cfg.MaxBatchRecords || bytes >= r.cfg.MaxBatchBytes { + if err := flushBatch(); err != nil { + return false, err + } + return true, nil + } + return false, nil + } flushUntil := func(target uint64) error { for { if target == 0 { @@ -675,13 +719,8 @@ func (r *recordStream) writerLoop() { if !ok { return r.streamError() } - batch = append(batch, req) - batches++ - bytes += len(req.Payload) - if batches >= r.cfg.MaxBatchRecords || bytes >= r.cfg.MaxBatchBytes { - if err := flushBatch(); err != nil { - return err - } + if _, err := appendBatch(req); err != nil { + return err } } } @@ -690,9 +729,15 @@ func (r *recordStream) writerLoop() { case <-r.ctx.Done(): return case req := <-r.sendCh: - batch = append(batch, req) - batches++ - bytes += len(req.Payload) + r.notifyOutboundReady() + flushed, err := appendBatch(req) + if err != nil { + r.abortRecord(err) + return + } + if flushed { + continue + } if len(batch) == 1 && r.cfg.MaxBatchDelay > 0 { if batchTimer == nil { batchTimer = time.NewTimer(r.cfg.MaxBatchDelay) @@ -701,36 +746,43 @@ func (r *recordStream) writerLoop() { } batchTimerCh = batchTimer.C } - if batches >= r.cfg.MaxBatchRecords || bytes >= r.cfg.MaxBatchBytes { - if err := flushBatch(); err != nil { - r.setTerminalError(err) - return - } - continue - } if ackSeq, sendNow := scheduleAck(len(batch) > 0, false); sendNow { if err := sendStandaloneAck(ackSeq); err != nil { - r.setTerminalError(err) + r.abortRecord(err) return } } case req := <-r.flushCh: + if req.ctx != nil && req.ctx.Err() != nil { + req.done <- req.ctx.Err() + continue + } err := flushUntil(req.targetSeq) if err == nil && req.forceAck { if ackSeq, sendNow := scheduleAck(len(batch) > 0, true); sendNow { err = sendStandaloneAck(ackSeq) } } + if err == nil && req.closeMode != recordCloseNone { + err = r.closeUnderlyingFromWriter(req) + } req.done <- err + if err != nil { + r.abortRecord(err) + return + } + if req.closeMode == recordCloseFull { + return + } case <-batchTimerCh: if err := flushBatch(); err != nil { - r.setTerminalError(err) + r.abortRecord(err) return } case <-r.ackCh: if ackSeq, sendNow := scheduleAck(len(batch) > 0, false); sendNow { if err := sendStandaloneAck(ackSeq); err != nil { - r.setTerminalError(err) + r.abortRecord(err) return } } @@ -738,7 +790,7 @@ func (r *recordStream) writerLoop() { stopAckTimer() if ackSeq, sendNow := scheduleAck(len(batch) > 0, true); sendNow { if err := sendStandaloneAck(ackSeq); err != nil { - r.setTerminalError(err) + r.abortRecord(err) return } } @@ -747,7 +799,7 @@ func (r *recordStream) writerLoop() { } func (r *recordStream) readLoop() { - defer close(r.recvCh) + defer r.closeReceive() defer close(r.readerCh) for { payload, err := readTransferFrame(r.stream) @@ -756,18 +808,12 @@ func (r *recordStream) readLoop() { r.markRemoteClosed(nil) return } - r.setReadError(err) + r.abortRecord(err) return } frame, err := decodeRecordFrame(payload) if err != nil { - _ = r.sendFailureFrame(RecordFailure{ - FailedSeq: r.nextInboundFailureSeq(), - Code: RecordErrorCodeProtocol, - Message: err.Error(), - }) - r.setReadError(err) - _ = r.stream.Reset(err) + r.abortProtocol(err) return } switch frame.Type { @@ -776,34 +822,31 @@ func (r *recordStream) readLoop() { if frame.AckSeq != 0 { r.obs.piggybackAckReceived.Add(1) if err := r.handleAckFrame(frame.AckSeq); err != nil { - r.setReadError(err) - _ = r.stream.Reset(err) + r.abortProtocol(err) return } } if err := r.handleBatchFrame(frame.Batch); err != nil { - _ = r.sendFailureFrame(RecordFailure{ - FailedSeq: r.nextInboundFailureSeq(), - Code: RecordErrorCodeProtocol, - Message: err.Error(), - }) - r.setReadError(err) - _ = r.stream.Reset(err) + r.abortProtocol(err) return } case recordFrameTypeAck: r.obs.ackFramesReceived.Add(1) if err := r.handleAckFrame(frame.AckSeq); err != nil { - r.setReadError(err) - _ = r.stream.Reset(err) + r.abortProtocol(err) return } case recordFrameTypeError: r.obs.errorFramesReceived.Add(1) - r.setReadError(frame.Failure) + r.abortRecord(frame.Failure) return + case recordFrameTypeFIN: + if err := r.receiveFIN(frame.FinalSeq); err != nil { + r.abortProtocol(err) + return + } default: - r.setReadError(errRecordFrameInvalid) + r.abortProtocol(errRecordFrameInvalid) return } } @@ -815,7 +858,7 @@ func (r *recordStream) handleBatchFrame(batch []recordOutboundMessage) error { } r.mu.Lock() expected := r.inboundReceivedSeq + 1 - if batch[0].Seq != expected { + if r.inboundClosed || batch[0].Seq != expected { r.mu.Unlock() return errRecordSeqInvalid } @@ -889,19 +932,6 @@ func (r *recordStream) markRemoteClosed(err error) { r.mu.Unlock() } -func (r *recordStream) setReadError(err error) { - if err == nil { - return - } - r.mu.Lock() - if r.readErr == nil { - r.readErr = err - } - r.signalStateLocked() - r.mu.Unlock() - r.cancel() -} - func (r *recordStream) setTerminalError(err error) { if err == nil { return @@ -918,38 +948,14 @@ func (r *recordStream) setTerminalError(err error) { r.cancel() } -func (r *recordStream) rollbackReservedOutboundLocked(seq uint64) { - if r == nil || seq == 0 { - return - } - if size, ok := r.outstandingSizes[seq]; ok { - delete(r.outstandingSizes, seq) - r.outstandingBytes -= size - if r.outstandingBytes < 0 { - r.outstandingBytes = 0 - } - r.outstandingRecords-- - if r.outstandingRecords < 0 { - r.outstandingRecords = 0 - } - } - if r.nextOutboundSeq == seq { - r.nextOutboundSeq-- - } - r.signalStateLocked() -} - func (r *recordStream) readError() error { if r == nil { return errRecordStreamNil } r.mu.Lock() defer r.mu.Unlock() - if r.readErr != nil { - return r.readErr - } - if r.terminalErr != nil { - return r.terminalErr + if err := r.streamErrorLocked(); err != nil { + return err } return io.EOF } @@ -970,6 +976,14 @@ func (r *recordStream) streamErrorLocked() error { if r.terminalErr != nil { return r.terminalErr } + if stream, ok := r.stream.(*streamHandle); ok && r.ctx.Err() != nil { + if err := stream.resetErrSnapshot(); err != nil { + return err + } + } + if r.ctx != nil { + return r.ctx.Err() + } return nil } @@ -1003,27 +1017,6 @@ func (r *recordStream) markAckSent(ackSeq uint64) { r.mu.Unlock() } -func (r *recordStream) flushAckNow() error { - if r == nil { - return errRecordStreamNil - } - req := recordFlushRequest{ - forceAck: true, - done: make(chan error, 1), - } - select { - case <-r.ctx.Done(): - return r.streamError() - case r.flushCh <- req: - } - select { - case <-r.ctx.Done(): - return r.streamError() - case err := <-req.done: - return err - } -} - func (r *recordStream) sendFailureFrame(failure RecordFailure) error { payload, err := encodeRecordErrorFrame(failure) if err != nil { @@ -1043,6 +1036,9 @@ func (r *recordStream) writePayloadFrame(payload []byte) error { if payload == nil { return nil } + if len(payload) > transferFrameMaxPayloadBytes { + return fmt.Errorf("%w: payload=%d max=%d", errTransferFrameTooLarge, len(payload), transferFrameMaxPayloadBytes) + } frame := buildTransferFrame(payload) r.writeMu.Lock() defer r.writeMu.Unlock() @@ -1067,10 +1063,18 @@ func (r *recordStream) nextOutboundForFlush() (recordOutboundMessage, bool) { case <-r.ctx.Done(): return recordOutboundMessage{}, false case req := <-r.sendCh: + r.notifyOutboundReady() return req, true } } +func (r *recordStream) notifyOutboundReady() { + select { + case r.sendReady <- struct{}{}: + default: + } +} + func (r *recordStream) nextInboundFailureSeq() uint64 { if r == nil { return 1 diff --git a/record_stream_test.go b/record_stream_test.go index a911e38..b8bf85e 100644 --- a/record_stream_test.go +++ b/record_stream_test.go @@ -166,6 +166,92 @@ func TestRecordStreamPropagatesStructuredFailure(t *testing.T) { } } +func TestRecordStreamPropagatesStructuredFailureUnderReceiveBackpressure(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(), 3*time.Second) + defer cancel() + local, err := client.OpenRecordStream(ctx, RecordOpenOptions{InboundQueueLimit: 1}) + if err != nil { + t.Fatal(err) + } + var remote RecordStream + select { + case remote = <-accepted: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + seq, err := local.WriteRecord(ctx, []byte("request")) + if err != nil { + t.Fatal(err) + } + if err := local.Flush(ctx); err != nil { + t.Fatal(err) + } + request, err := remote.ReadRecord(ctx) + if err != nil { + t.Fatal(err) + } + for i := 0; i < 2; i++ { + if _, err := remote.WriteRecord(ctx, []byte("pending response")); err != nil { + t.Fatal(err) + } + } + if err := remote.Flush(ctx); err != nil { + t.Fatal(err) + } + record := local.(*recordStream) + for { + record.mu.Lock() + received := record.inboundReceivedSeq + record.mu.Unlock() + if received == 2 && len(record.recvCh) == 1 { + break + } + select { + case <-ctx.Done(): + t.Fatal("record receive queue did not fill") + case <-time.After(time.Millisecond): + } + } + + want := RecordFailure{FailedSeq: request.Seq, Code: "disk_full", Retryable: true, Message: "disk full"} + if err := remote.FailRecord(request.Seq, want); err != nil { + t.Fatalf("FailRecord failed: %v", err) + } + select { + case <-local.Context().Done(): + case <-ctx.Done(): + t.Fatal("remote failure did not terminate local record stream") + } + _, err = local.BarrierTo(ctx, seq) + var got RecordFailure + if !errors.As(err, &got) || got != want { + t.Fatalf("structured failure lost under receive backpressure: got=%T %v, want=%+v", err, err, want) + } +} + func TestRecordStreamBackpressureUsesUnackedRecords(t *testing.T) { server := NewServer().(*ServerCommon) secret := []byte("0123456789abcdef0123456789abcdef") @@ -340,7 +426,9 @@ func TestRecordStreamConcurrentWritesStayOrdered(t *testing.T) { server.SetSecretKey(secret) }) - const total = 64 + const total = 512 + writeCtx, cancelWrites := context.WithTimeout(context.Background(), 5*time.Second) + defer cancelWrites() receivedCh := make(chan RecordMessage, total) handlerDone := make(chan error, 1) server.SetRecordStreamHandler(func(info RecordAcceptInfo) error { @@ -387,14 +475,14 @@ func TestRecordStreamConcurrentWritesStayOrdered(t *testing.T) { go func() { defer wg.Done() payload := []byte("item-" + strconv.Itoa(index)) - if _, err := stream.WriteRecord(context.Background(), payload); err != nil { + if _, err := stream.WriteRecord(writeCtx, payload); err != nil { t.Errorf("WriteRecord(%d) failed: %v", index, err) } }() } wg.Wait() - if acked, err := stream.Barrier(context.Background()); err != nil { + if acked, err := stream.Barrier(writeCtx); err != nil { t.Fatalf("Barrier failed: %v", err) } else if got, want := acked, uint64(total); got != want { t.Fatalf("Barrier acked=%d want=%d", got, want) diff --git a/review_fix_regression_test.go b/review_fix_regression_test.go new file mode 100644 index 0000000..0576913 --- /dev/null +++ b/review_fix_regression_test.go @@ -0,0 +1,197 @@ +package notify + +import ( + "b612.me/stario" + "context" + "errors" + netpkg "net" + "sync/atomic" + "testing" + "time" +) + +type gatedDedicatedReadConn struct { + netpkg.Conn + total int + read int + reached chan struct{} + release chan struct{} + closed bool +} + +func (c *gatedDedicatedReadConn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p) + c.read += n + if !c.closed && c.read >= c.total { + c.closed = true + close(c.reached) + <-c.release + } + return n, err +} + +func TestControlMessageErrorPreservesTransportDetachedSentinel(t *testing.T) { + for name, decode := range map[string]func(string) error{ + "bulk": bulkControlMessageError, + "stream": streamControlMessageError, + } { + t.Run(name, func(t *testing.T) { + err := decode("transport detached: stale transport generation=7") + if !errors.Is(err, errTransportDetached) { + t.Fatalf("decoded error = %v, want transport detached sentinel", err) + } + }) + } +} + +func TestSendDedicatedBulkAttachRequestRejectsNilDialResult(t *testing.T) { + client := &ClientCommon{} + bulk := newBulkHandle(context.Background(), newBulkRuntime("cblk"), clientFileScope(), BulkOpenRequest{ + BulkID: "nil-dial", + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + if _, err := client.sendDedicatedBulkAttachRequest(context.Background(), nil, bulk); !errors.Is(err, errTransportDetached) { + t.Fatalf("nil dial result error = %v, want transport detached", err) + } +} + +func TestDedicatedSidecarCloseWaitsForAttachment(t *testing.T) { + left, right := netpkg.Pipe() + defer left.Close() + defer right.Close() + sidecar := newBulkDedicatedSidecar(left, 1) + entered := make(chan struct{}) + release := make(chan struct{}) + attachDone := make(chan error, 1) + go func() { + attachDone <- sidecar.withConn(func(conn netpkg.Conn) error { + close(entered) + <-release + return nil + }) + }() + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("sidecar attachment did not start") + } + closeDone := make(chan struct{}) + go func() { + sidecar.close() + close(closeDone) + }() + select { + case <-closeDone: + t.Fatal("sidecar closed while attachment still held") + case <-time.After(25 * time.Millisecond): + } + close(release) + if err := <-attachDone; err != nil { + t.Fatalf("sidecar attachment error = %v", err) + } + select { + case <-closeDone: + case <-time.After(time.Second): + t.Fatal("sidecar close did not finish after attachment release") + } +} + +func TestBulkAcceptDispatchRejectsResetBeforeHandler(t *testing.T) { + bulk := newBulkHandle(context.Background(), newBulkRuntime("sblk"), serverFileDomain+":test", BulkOpenRequest{ + BulkID: "stale-dispatch", + }, 0, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.markReset(errTransportDetached) + var calls atomic.Int32 + err := dispatchBulkAccept(func(BulkAcceptInfo) error { + calls.Add(1) + return nil + }, bulk, BulkAcceptInfo{Bulk: bulk}) + if !errors.Is(err, errTransportDetached) { + t.Fatalf("stale dispatch error = %v, want transport detached", err) + } + if got := calls.Load(); got != 0 { + t.Fatalf("stale dispatch handler calls = %d, want 0", got) + } +} + +func TestClientDedicatedSidecarDropsFrameAfterTransportReattach(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, ^uint32(0)) + oldLeft, oldRight := netpkg.Pipe() + defer oldRight.Close() + epoch := client.beginClientSessionEpoch() + client.setClientSessionRuntime(newClientSessionRuntime(oldLeft, stopCtx, stopFn, queue, epoch)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + oldRoute := client.clientSessionRouteSnapshot() + + sidecarLeft, sidecarRight := netpkg.Pipe() + defer sidecarRight.Close() + payload, err := client.encodeDedicatedBulkBatchPayload(101, []bulkDedicatedSendRequest{{ + Type: bulkFastPayloadTypeData, + Seq: 1, + Payload: []byte("stale"), + }}) + if err != nil { + t.Fatalf("encode stale sidecar payload: %v", err) + } + frameRead := make(chan struct{}) + releaseFrame := make(chan struct{}) + readConn := &gatedDedicatedReadConn{ + Conn: sidecarLeft, + total: bulkDedicatedRecordHeaderLen + len(payload), + reached: frameRead, + release: releaseFrame, + } + sidecar := newBulkDedicatedSidecar(readConn, 1) + loopDone := make(chan struct{}) + go func() { + client.readDedicatedSidecarLoopAtRoute(sidecar, oldRoute) + close(loopDone) + }() + + writeDone := make(chan error, 1) + go func() { writeDone <- writeBulkDedicatedRecord(sidecarRight, payload) }() + select { + case <-frameRead: + case <-time.After(time.Second): + t.Fatal("sidecar loop did not read the complete frame") + } + + newLeft, newRight := netpkg.Pipe() + defer newRight.Close() + if err := client.attachClientSessionTransport(newLeft); err != nil { + t.Fatalf("attach replacement client transport: %v", err) + } + + runtime := client.getBulkRuntime() + bulk := newBulkHandle(stopCtx, runtime, clientFileScope(), BulkOpenRequest{ + BulkID: "stale-sidecar-frame", + DataID: 101, + Range: BulkRange{Length: 32}, + }, oldRoute.epoch, nil, nil, 0, nil, nil, nil, nil, nil) + bulk.setClientSnapshotOwner(client) + bulk.setClientSessionRoute(oldRoute) + if err := runtime.registerInbound(clientFileScope(), bulk); err != nil { + t.Fatalf("register stale-route bulk: %v", err) + } + defer bulk.markReset(errors.New("test cleanup")) + + close(releaseFrame) + if err := <-writeDone; err != nil { + t.Fatalf("write stale sidecar payload: %v", err) + } + select { + case <-loopDone: + case <-time.After(time.Second): + t.Fatal("stale-route sidecar loop did not stop after reattach") + } + + bulk.mu.Lock() + defer bulk.mu.Unlock() + if len(bulk.readQueue) != 0 || len(bulk.readBuf.data) != 0 || bulk.resetErr != nil { + t.Fatalf("stale sidecar frame mutated bulk: queued=%d buffered=%d reset=%v", len(bulk.readQueue), len(bulk.readBuf.data), bulk.resetErr) + } +} diff --git a/server.go b/server.go index b16dee9..351dd08 100644 --- a/server.go +++ b/server.go @@ -63,6 +63,8 @@ type ServerCommon struct { streamRuntime *streamRuntime recordRuntime *recordRuntime bulkRuntime *bulkRuntime + bulkRecovery *bulkRecoveryQueue + bulkRecoveryMu sync.Mutex bulkOpenTuning BulkOpenTuning udpWriteGateOnce sync.Once udpWriteGate chan struct{} @@ -103,6 +105,7 @@ func NewServer() Server { server.streamRuntime = newStreamRuntime("sstrm") server.recordRuntime = newRecordRuntime() server.bulkRuntime = newBulkRuntime("sblk") + server.bulkRecovery = newBulkRecoveryQueue(server.reportBulkRecoveryError) server.bulkOpenTuning = defaultBulkOpenTuning() server.bulkDedicatedSidecars = make(map[*LogicalConn]map[uint32]*bulkDedicatedSidecar) server.connectionRetryState = newConnectionRetryState() diff --git a/server_bulk.go b/server_bulk.go index 1b2b196..7898b27 100644 --- a/server_bulk.go +++ b/server_bulk.go @@ -18,20 +18,20 @@ func (s *ServerCommon) OpenBulkLogical(ctx context.Context, logical *LogicalConn switch opt.Mode { case BulkOpenModeDedicated: opt.Dedicated = true - return s.openBulkLogicalWithMode(ctx, logical, opt) + return s.openBulkLogicalWithMode(ctx, logical, opt, false) case BulkOpenModeAuto: if err := logicalDedicatedBulkSupportError(logical); err == nil { dedicatedOpt := opt dedicatedOpt.Mode = BulkOpenModeDedicated dedicatedOpt.Dedicated = true - bulk, dedicatedErr := s.openBulkLogicalWithMode(ctx, logical, dedicatedOpt) + bulk, dedicatedErr := s.openBulkLogicalWithMode(ctx, logical, dedicatedOpt, true) if dedicatedErr == nil { return bulk, nil } sharedOpt := opt sharedOpt.Mode = BulkOpenModeShared sharedOpt.Dedicated = false - sharedBulk, sharedErr := s.openBulkLogicalWithMode(ctx, logical, sharedOpt) + sharedBulk, sharedErr := s.openBulkLogicalWithMode(ctx, logical, sharedOpt, false) if sharedErr == nil { return sharedBulk, nil } @@ -39,19 +39,19 @@ func (s *ServerCommon) OpenBulkLogical(ctx context.Context, logical *LogicalConn } opt.Mode = BulkOpenModeShared opt.Dedicated = false - return s.openBulkLogicalWithMode(ctx, logical, opt) + return s.openBulkLogicalWithMode(ctx, logical, opt, false) case BulkOpenModeShared, BulkOpenModeDefault: opt.Mode = BulkOpenModeShared opt.Dedicated = false - return s.openBulkLogicalWithMode(ctx, logical, opt) + return s.openBulkLogicalWithMode(ctx, logical, opt, false) default: opt.Mode = BulkOpenModeShared opt.Dedicated = false - return s.openBulkLogicalWithMode(ctx, logical, opt) + return s.openBulkLogicalWithMode(ctx, logical, opt, false) } } -func (s *ServerCommon) openBulkLogicalWithMode(ctx context.Context, logical *LogicalConn, opt BulkOpenOptions) (Bulk, error) { +func (s *ServerCommon) openBulkLogicalWithMode(ctx context.Context, logical *LogicalConn, opt BulkOpenOptions, waitForReset bool) (Bulk, error) { if s == nil { return nil, errBulkServerNil } @@ -59,6 +59,7 @@ func (s *ServerCommon) openBulkLogicalWithMode(ctx context.Context, logical *Log if logical == nil { return nil, errBulkLogicalConnNil } + transport := logical.CurrentTransportConn() runtime := s.getBulkRuntime() if runtime == nil { return nil, errBulkRuntimeNil @@ -76,83 +77,111 @@ func (s *ServerCommon) openBulkLogicalWithMode(ctx context.Context, logical *Log if _, exists := runtime.lookup(scope, req.BulkID); exists { return nil, errBulkAlreadyExists } - if req.Dedicated { - if req.DataID == 0 { - req.DataID = runtime.nextDataID() + if req.DataID == 0 { + var reserveErr error + req.DataID, reserveErr = runtime.reserveDataID(scope, 0) + if reserveErr != nil { + return nil, reserveErr } + } + if req.Dedicated { if req.AttachToken == "" { req.AttachToken = newBulkAttachToken() } - bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, logical.CurrentTransportConn(), logical.transportGenerationSnapshot(), serverBulkCloseSender(s, logical, nil), serverBulkResetSender(s, logical, nil), serverBulkDataSender(s, logical.CurrentTransportConn()), serverBulkWriteSender(s, logical, logical.CurrentTransportConn()), serverBulkReleaseSender(s, logical, logical.CurrentTransportConn())) + bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, logical.transportGenerationSnapshot(), serverBulkCloseSender(s, logical, transport), serverBulkResetSender(s, logical, transport), serverBulkDataSender(s, transport), serverBulkWriteSender(s, logical, transport), serverBulkReleaseSender(s, logical, transport)) bulk.markAcceptHandled() - if err := runtime.register(scope, bulk); err != nil { + if err := runtime.adoptReserved(scope, bulk); err != nil { + runtime.releaseDataID(scope, req.DataID) return nil, err } s.attachServerDedicatedSidecarIfExists(logical, bulk) - resp, err := sendBulkOpenServerLogical(ctx, s, logical, req) + resp, err := sendBulkOpenServerTransport(ctx, s, transport, req) if err != nil { + runtime.releaseDataID(scope, req.DataID) + cleanupErr := s.cleanupBulkLogicalReset(ctx, logical, transport, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: err.Error()}, waitForReset) bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } if resp.DataID != 0 && resp.DataID != req.DataID { err = errBulkAlreadyExists - _, _ = sendBulkResetServerLogical(context.Background(), s, logical, BulkResetRequest{ + cleanupErr := s.cleanupBulkLogicalReset(ctx, logical, transport, BulkResetRequest{ BulkID: req.BulkID, - DataID: req.DataID, Error: "bulk dedicated data id mismatch", - }) + }, waitForReset) bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } if resp.TransportGeneration != 0 { - bulk.transportGeneration = resp.TransportGeneration + bulk.setTransportGeneration(resp.TransportGeneration) } if resp.FastPathVersion != 0 { - bulk.fastPathVersion = normalizeBulkFastPathVersion(resp.FastPathVersion) + bulk.setFastPathVersion(resp.FastPathVersion) } if resp.AttachToken != "" { bulk.setDedicatedAttachToken(resp.AttachToken) } if err := bulk.waitAcceptReady(ctx); err != nil { - _, _ = sendBulkResetServerLogical(context.Background(), s, logical, BulkResetRequest{ - BulkID: req.BulkID, - DataID: req.DataID, - Error: err.Error(), - }) + var cleanupErr error + if bulk.resetErrSnapshot() == nil { + cleanupErr = s.cleanupBulkLogicalReset(ctx, logical, transport, BulkResetRequest{ + BulkID: req.BulkID, + DataID: req.DataID, + Error: err.Error(), + }, waitForReset) + } else { + s.bestEffortBulkResetLogical(logical, transport, BulkResetRequest{ + BulkID: req.BulkID, + DataID: req.DataID, + Error: err.Error(), + }) + } bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } return bulk, nil } - resp, err := sendBulkOpenServerLogical(ctx, s, logical, req) - if err != nil { + bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, logical.transportGenerationSnapshot(), serverBulkCloseSender(s, logical, transport), serverBulkResetSender(s, logical, transport), serverBulkDataSender(s, transport), serverBulkWriteSender(s, logical, transport), serverBulkReleaseSender(s, logical, transport)) + bulk.markAcceptHandled() + if err := runtime.adoptReserved(scope, bulk); err != nil { + runtime.releaseDataID(scope, req.DataID) return nil, err } - if resp.DataID != 0 { - req.DataID = resp.DataID + resp, err := sendBulkOpenServerTransport(ctx, s, transport, req) + if err != nil { + s.bestEffortBulkResetLogical(logical, transport, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: err.Error()}) + bulk.markReset(err) + return nil, err + } + if resp.DataID != 0 && resp.DataID != req.DataID { + err = errBulkAlreadyExists + s.bestEffortBulkResetLogical(logical, transport, BulkResetRequest{BulkID: req.BulkID, Error: "bulk data id mismatch"}) + bulk.markReset(err) + return nil, err } if resp.FastPathVersion != 0 { - req.FastPathVersion = resp.FastPathVersion + bulk.setFastPathVersion(resp.FastPathVersion) } - req.Dedicated = resp.Dedicated - if resp.AttachToken != "" { - req.AttachToken = resp.AttachToken - } - if req.DataID == 0 { - return nil, errBulkDataIDEmpty - } - transport := logical.CurrentTransportConn() - bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, resp.TransportGeneration, serverBulkCloseSender(s, logical, nil), serverBulkResetSender(s, logical, nil), serverBulkDataSender(s, transport), serverBulkWriteSender(s, logical, transport), serverBulkReleaseSender(s, logical, transport)) - bulk.markAcceptHandled() - if err := runtime.register(scope, bulk); err != nil { - _, _ = sendBulkResetServerLogical(context.Background(), s, logical, BulkResetRequest{ - BulkID: req.BulkID, - DataID: req.DataID, - Error: err.Error(), - }) + if resp.Dedicated { + err = errBulkRejected + s.bestEffortBulkResetLogical(logical, transport, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: "shared bulk upgraded to dedicated"}) + bulk.markReset(err) return nil, err } - s.attachServerDedicatedSidecarIfExists(logical, bulk) + if resp.AttachToken != "" { + bulk.setDedicatedAttachToken(resp.AttachToken) + } + if resp.TransportGeneration != 0 { + bulk.setTransportGeneration(resp.TransportGeneration) + } return bulk, nil } @@ -161,20 +190,20 @@ func (s *ServerCommon) OpenBulkTransport(ctx context.Context, transport *Transpo switch opt.Mode { case BulkOpenModeDedicated: opt.Dedicated = true - return s.openBulkTransportWithMode(ctx, transport, opt) + return s.openBulkTransportWithMode(ctx, transport, opt, false) case BulkOpenModeAuto: if err := transportDedicatedBulkSupportError(transport); err == nil { dedicatedOpt := opt dedicatedOpt.Mode = BulkOpenModeDedicated dedicatedOpt.Dedicated = true - bulk, dedicatedErr := s.openBulkTransportWithMode(ctx, transport, dedicatedOpt) + bulk, dedicatedErr := s.openBulkTransportWithMode(ctx, transport, dedicatedOpt, true) if dedicatedErr == nil { return bulk, nil } sharedOpt := opt sharedOpt.Mode = BulkOpenModeShared sharedOpt.Dedicated = false - sharedBulk, sharedErr := s.openBulkTransportWithMode(ctx, transport, sharedOpt) + sharedBulk, sharedErr := s.openBulkTransportWithMode(ctx, transport, sharedOpt, false) if sharedErr == nil { return sharedBulk, nil } @@ -182,19 +211,19 @@ func (s *ServerCommon) OpenBulkTransport(ctx context.Context, transport *Transpo } opt.Mode = BulkOpenModeShared opt.Dedicated = false - return s.openBulkTransportWithMode(ctx, transport, opt) + return s.openBulkTransportWithMode(ctx, transport, opt, false) case BulkOpenModeShared, BulkOpenModeDefault: opt.Mode = BulkOpenModeShared opt.Dedicated = false - return s.openBulkTransportWithMode(ctx, transport, opt) + return s.openBulkTransportWithMode(ctx, transport, opt, false) default: opt.Mode = BulkOpenModeShared opt.Dedicated = false - return s.openBulkTransportWithMode(ctx, transport, opt) + return s.openBulkTransportWithMode(ctx, transport, opt, false) } } -func (s *ServerCommon) openBulkTransportWithMode(ctx context.Context, transport *TransportConn, opt BulkOpenOptions) (Bulk, error) { +func (s *ServerCommon) openBulkTransportWithMode(ctx context.Context, transport *TransportConn, opt BulkOpenOptions, waitForReset bool) (Bulk, error) { if s == nil { return nil, errBulkServerNil } @@ -223,85 +252,164 @@ func (s *ServerCommon) openBulkTransportWithMode(ctx context.Context, transport if _, exists := runtime.lookup(scope, req.BulkID); exists { return nil, errBulkAlreadyExists } - if req.Dedicated { - if req.DataID == 0 { - req.DataID = runtime.nextDataID() + if req.DataID == 0 { + var reserveErr error + req.DataID, reserveErr = runtime.reserveDataID(scope, 0) + if reserveErr != nil { + return nil, reserveErr } + } + if req.Dedicated { if req.AttachToken == "" { req.AttachToken = newBulkAttachToken() } bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, transport.TransportGeneration(), serverBulkCloseSender(s, logical, transport), serverBulkResetSender(s, logical, transport), serverBulkDataSender(s, transport), serverBulkWriteSender(s, logical, transport), serverBulkReleaseSender(s, logical, transport)) bulk.markAcceptHandled() - if err := runtime.register(scope, bulk); err != nil { + if err := runtime.adoptReserved(scope, bulk); err != nil { + runtime.releaseDataID(scope, req.DataID) return nil, err } s.attachServerDedicatedSidecarIfExists(logical, bulk) resp, err := sendBulkOpenServerTransport(ctx, s, transport, req) if err != nil { + runtime.releaseDataID(scope, req.DataID) + cleanupErr := s.cleanupBulkTransportReset(ctx, transport, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: err.Error()}, waitForReset) bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } if resp.DataID != 0 && resp.DataID != req.DataID { err = errBulkAlreadyExists - _, _ = sendBulkResetServerTransport(context.Background(), s, transport, BulkResetRequest{ + cleanupErr := s.cleanupBulkTransportReset(ctx, transport, BulkResetRequest{ BulkID: req.BulkID, - DataID: req.DataID, Error: "bulk dedicated data id mismatch", - }) + }, waitForReset) bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } if resp.TransportGeneration != 0 { - bulk.transportGeneration = resp.TransportGeneration + bulk.setTransportGeneration(resp.TransportGeneration) } if resp.FastPathVersion != 0 { - bulk.fastPathVersion = normalizeBulkFastPathVersion(resp.FastPathVersion) + bulk.setFastPathVersion(resp.FastPathVersion) } if resp.AttachToken != "" { bulk.setDedicatedAttachToken(resp.AttachToken) } if err := bulk.waitAcceptReady(ctx); err != nil { - _, _ = sendBulkResetServerTransport(context.Background(), s, transport, BulkResetRequest{ - BulkID: req.BulkID, - DataID: req.DataID, - Error: err.Error(), - }) + var cleanupErr error + if bulk.resetErrSnapshot() == nil { + cleanupErr = s.cleanupBulkTransportReset(ctx, transport, BulkResetRequest{ + BulkID: req.BulkID, + DataID: req.DataID, + Error: err.Error(), + }, waitForReset) + } else { + s.bestEffortBulkResetTransport(transport, BulkResetRequest{ + BulkID: req.BulkID, + DataID: req.DataID, + Error: err.Error(), + }) + } bulk.markReset(err) + if cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } return nil, err } return bulk, nil } + bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, transport.TransportGeneration(), serverBulkCloseSender(s, logical, transport), serverBulkResetSender(s, logical, transport), serverBulkDataSender(s, transport), serverBulkWriteSender(s, logical, transport), serverBulkReleaseSender(s, logical, transport)) + bulk.markAcceptHandled() + if err := runtime.adoptReserved(scope, bulk); err != nil { + runtime.releaseDataID(scope, req.DataID) + return nil, err + } resp, err := sendBulkOpenServerTransport(ctx, s, transport, req) if err != nil { + s.bestEffortBulkResetTransport(transport, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: err.Error()}) + bulk.markReset(err) return nil, err } - if resp.DataID != 0 { - req.DataID = resp.DataID + if resp.DataID != 0 && resp.DataID != req.DataID { + err = errBulkAlreadyExists + s.bestEffortBulkResetTransport(transport, BulkResetRequest{BulkID: req.BulkID, Error: "bulk data id mismatch"}) + bulk.markReset(err) + return nil, err } if resp.FastPathVersion != 0 { - req.FastPathVersion = resp.FastPathVersion + bulk.setFastPathVersion(resp.FastPathVersion) } - req.Dedicated = resp.Dedicated - if resp.AttachToken != "" { - req.AttachToken = resp.AttachToken - } - if req.DataID == 0 { - return nil, errBulkDataIDEmpty - } - bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, resp.TransportGeneration, serverBulkCloseSender(s, logical, transport), serverBulkResetSender(s, logical, transport), serverBulkDataSender(s, transport), serverBulkWriteSender(s, logical, transport), serverBulkReleaseSender(s, logical, transport)) - bulk.markAcceptHandled() - if err := runtime.register(scope, bulk); err != nil { - _, _ = sendBulkResetServerTransport(context.Background(), s, transport, BulkResetRequest{ - BulkID: req.BulkID, - DataID: req.DataID, - Error: err.Error(), - }) + if resp.Dedicated { + err = errBulkRejected + s.bestEffortBulkResetTransport(transport, BulkResetRequest{BulkID: req.BulkID, DataID: req.DataID, Error: "shared bulk upgraded to dedicated"}) + bulk.markReset(err) return nil, err } - s.attachServerDedicatedSidecarIfExists(logical, bulk) + if resp.AttachToken != "" { + bulk.setDedicatedAttachToken(resp.AttachToken) + } + if resp.TransportGeneration != 0 { + bulk.setTransportGeneration(resp.TransportGeneration) + } return bulk, nil } +func (s *ServerCommon) bestEffortBulkResetLogical(logical *LogicalConn, transport *TransportConn, req BulkResetRequest) { + if s == nil || logical == nil { + return + } + task := newServerBulkResetRecoveryTask(s, logical, transport, req) + q := s.bulkRecoveryQueue() + if !q.enqueue(task) { + s.handleBulkRecoveryOverflow(logical, transport, req) + } +} + +func (s *ServerCommon) bulkRecoveryQueue() *bulkRecoveryQueue { + if s == nil { + return nil + } + s.bulkRecoveryMu.Lock() + defer s.bulkRecoveryMu.Unlock() + if s.bulkRecovery == nil { + s.bulkRecovery = newBulkRecoveryQueue(s.reportBulkRecoveryError) + } + return s.bulkRecovery +} + +func (s *ServerCommon) bestEffortBulkResetTransport(transport *TransportConn, req BulkResetRequest) { + if s == nil || transport == nil { + return + } + task := newServerBulkResetRecoveryTask(s, transport.logicalConnSnapshot(), transport, req) + q := s.bulkRecoveryQueue() + if !q.enqueue(task) { + s.handleBulkRecoveryOverflow(transport.logicalConnSnapshot(), transport, req) + } +} + +func newServerBulkResetRecoveryTask(s *ServerCommon, logical *LogicalConn, transport *TransportConn, req BulkResetRequest) bulkRecoveryTask { + return func(ctx context.Context) error { + if transport == nil { + if logical != nil { + return transportDetachedErrorForLogical(logical) + } + return errTransportDetached + } + _, err := sendBulkResetServerTransport(ctx, s, transport, req) + if errors.Is(err, errBulkNotFound) { + return nil + } + return err + } +} + func serverBulkRequest(runtime *bulkRuntime, opt BulkOpenOptions) BulkOpenRequest { opt = normalizeBulkOpenOptions(opt) id := opt.ID @@ -332,13 +440,14 @@ func serverBulkCloseSender(s *ServerCommon, logical *LogicalConn, transport *Tra } req := BulkCloseRequest{ BulkID: bulk.ID(), + DataID: bulk.dataIDSnapshot(), Full: full, } - if logical != nil { - _, err := sendBulkCloseServerLogical(ctx, s, logical, req) + if transport != nil { + _, err := sendBulkCloseServerTransport(ctx, s, transport, req) return err } - _, err := sendBulkCloseServerTransport(ctx, s, transport, req) + _, err := sendBulkCloseServerLogical(ctx, s, logical, req) return err } } @@ -356,11 +465,11 @@ func serverBulkResetSender(s *ServerCommon, logical *LogicalConn, transport *Tra DataID: bulk.dataIDSnapshot(), Error: message, } - if logical != nil { - _, err := sendBulkResetServerLogical(ctx, s, logical, req) + if transport != nil { + _, err := sendBulkResetServerTransport(ctx, s, transport, req) return err } - _, err := sendBulkResetServerTransport(ctx, s, transport, req) + _, err := sendBulkResetServerLogical(ctx, s, logical, req) return err } } @@ -461,7 +570,7 @@ func serverBulkReleaseSender(s *ServerCommon, logical *LogicalConn, transport *T Bytes: bytes, Chunks: chunks, } - if transport != nil && transport.IsCurrent() { + if transport != nil { return sendBulkReleaseServerTransport(ctx, s, transport, req) } return sendBulkReleaseServerLogical(ctx, s, logical, req) diff --git a/server_inbound_source.go b/server_inbound_source.go index e5aa63b..406b03a 100644 --- a/server_inbound_source.go +++ b/server_inbound_source.go @@ -84,6 +84,15 @@ func (s *ServerCommon) pushTransportPayloadSourceFast(payload []byte, release fu } return false } + if err := validateTransportFramePayloadLen(payload); err != nil { + if release != nil { + release() + } + if s.showError || s.debugMode { + fmt.Println("server enqueue inbound frame error", err) + } + return true + } frame := queue.BuildMessage(payload) if release != nil { release() @@ -117,7 +126,7 @@ func (s *ServerCommon) pushTransportPayloadSourceFast(payload []byte, release fu plainRelease() } s.wg.Add(1) - if !dispatcher.Dispatch(serverInboundDispatchSource(source), func() { + if !dispatcher.DispatchSized(serverInboundDispatchSource(source), len(owned), func() { defer s.wg.Done() now := time.Now() if err := s.dispatchInboundTransportPlain(logical, transport, inboundConn, owned, now); err != nil && (s.showError || s.debugMode) { diff --git a/server_listen.go b/server_listen.go index db92a8c..f781bed 100644 --- a/server_listen.go +++ b/server_listen.go @@ -6,7 +6,6 @@ import ( "context" "errors" "fmt" - "math" "math/rand" "net" "os" @@ -30,7 +29,7 @@ func (s *ServerCommon) Listen(network string, addr string) error { } s.applySignalReliabilityTransportDefault(transport.IsUDPNetwork(network)) stopCtx, stopFn := context.WithCancel(context.Background()) - queue := stario.NewQueueCtx(stopCtx, 128, math.MaxUint32) + queue := stario.NewQueueCtx(stopCtx, 128, transportFrameMaxPayloadBytes) s.setServerSessionRuntime(&serverSessionRuntime{ stopCtx: stopCtx, stopFn: stopFn, @@ -70,7 +69,7 @@ func (s *ServerCommon) ListenByListener(listener net.Listener) error { } s.applySignalReliabilityTransportDefault(false) stopCtx, stopFn := context.WithCancel(context.Background()) - queue := stario.NewQueueCtx(stopCtx, 128, math.MaxUint32) + queue := stario.NewQueueCtx(stopCtx, 128, transportFrameMaxPayloadBytes) s.setServerSessionRuntime(&serverSessionRuntime{ stopCtx: stopCtx, stopFn: stopFn, @@ -262,7 +261,7 @@ func (s *ServerCommon) loadMessageLoop(logicalStopCtx context.Context, transport } msg := data s.wg.Add(1) - if !dispatcher.Dispatch(serverInboundDispatchSource(msg.Conn), func() { + if !dispatcher.DispatchSized(serverInboundDispatchSource(msg.Conn), len(msg.Msg), func() { defer s.wg.Done() logical, transport := s.resolveInboundSource(msg.Conn) if logical == nil { diff --git a/server_send.go b/server_send.go index 62898cd..710bc46 100644 --- a/server_send.go +++ b/server_send.go @@ -453,7 +453,7 @@ func (s *ServerCommon) writeControlEnvelopePayload(logical *LogicalConn, transpo if s.serverUDPListenerSnapshot() != nil { return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, writeTimeout) } - binding := logical.transportBindingSnapshot() + binding := serverTransportBindingSnapshotForConn(logical, transport, conn) if binding == nil || binding.queueSnapshot() == nil { return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, writeTimeout) } @@ -530,6 +530,9 @@ func (s *ServerCommon) writeEnvelopePayloadContextTimeout(ctx context.Context, l if transport == nil || transport.RemoteAddr() == nil { return transportDetachedErrorForTransport(transport) } + if err := validateTransportFramePayloadLen(payload); err != nil { + return err + } data := queue.BuildMessage(payload) deadline := earlierWriteDeadline(writeDeadlineFromTimeout(writeTimeout), contextDeadline(ctx)) return s.withUDPWriteLockDeadline(ctx, deadline, func() error { @@ -543,10 +546,7 @@ func (s *ServerCommon) writeEnvelopePayloadContextTimeout(ctx context.Context, l return err }) } - var binding *transportBinding - if logical != nil { - binding = logical.transportBindingSnapshot() - } + binding := serverTransportBindingSnapshotForConn(logical, transport, conn) if conn == nil { if binding == nil { return os.ErrClosed diff --git a/server_session.go b/server_session.go index f20ac48..199812b 100644 --- a/server_session.go +++ b/server_session.go @@ -62,6 +62,15 @@ func (s *ServerCommon) detachClientSessionTransport(client *ClientConn, reason s } func (s *ServerCommon) detachLogicalSessionTransport(logical *LogicalConn, reason string, err error) { + if s == nil || logical == nil { + return + } + logical.transportLifecycleMu.Lock() + defer logical.transportLifecycleMu.Unlock() + s.detachLogicalSessionTransportLocked(logical, reason, err) +} + +func (s *ServerCommon) detachLogicalSessionTransportLocked(logical *LogicalConn, reason string, err error) { if s == nil || logical == nil { return } @@ -185,6 +194,14 @@ func (s *ServerCommon) attachAcceptedLogicalTransport(logical *LogicalConn, addr return errors.New("logical conn is nil") } logical.setServer(s) + logical.transportLifecycleMu.Lock() + defer logical.transportLifecycleMu.Unlock() + if oldConn := logical.transportSnapshot(); tuConn != nil && oldConn != nil && oldConn != tuConn { + // Reusing a logical peer is a transport-generation boundary. Retire every + // operation pinned to the old connection before publishing the replacement; + // otherwise an in-flight dedicated attach can revive an old bulk. + s.detachLogicalSessionTransportLocked(logical, "server transport replaced", transportDetachedError("server transport replaced", nil)) + } return logical.attachAcceptedTransport(addr, tuConn) } diff --git a/server_stream.go b/server_stream.go index c053e2c..dd133cc 100644 --- a/server_stream.go +++ b/server_stream.go @@ -17,38 +17,7 @@ func (s *ServerCommon) OpenStreamLogical(ctx context.Context, logical *LogicalCo if logical == nil { return nil, errStreamLogicalConnNil } - runtime := s.getStreamRuntime() - if runtime == nil { - return nil, errStreamRuntimeNil - } - req := serverStreamRequest(runtime, opt) - scope := serverFileScope(logical) - if _, exists := runtime.lookup(scope, req.StreamID); exists { - return nil, errStreamAlreadyExists - } - resp, err := sendStreamOpenServerLogical(ctx, s, logical, req) - if err != nil { - return nil, err - } - if resp.DataID != 0 { - req.DataID = resp.DataID - } - if resp.FastPathVersion != 0 { - req.FastPathVersion = resp.FastPathVersion - } else { - req.FastPathVersion = streamFastPathVersionV1 - } - req.Metadata = mergeStreamMetadata(req.Metadata, resp.Metadata) - transport := logical.CurrentTransportConn() - stream := newStreamHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, resp.TransportGeneration, serverStreamCloseSender(s, logical, nil), serverStreamResetSender(s, logical, nil), serverStreamDataSender(s, transport), runtime.configSnapshot()) - if err := runtime.register(scope, stream); err != nil { - _, _ = sendStreamResetServerLogical(context.Background(), s, logical, StreamResetRequest{ - StreamID: req.StreamID, - Error: err.Error(), - }) - return nil, err - } - return stream, nil + return s.openStreamTransport(ctx, logical.CurrentTransportConn(), opt) } func (s *ServerCommon) OpenStreamTransport(ctx context.Context, transport *TransportConn, opt StreamOpenOptions) (Stream, error) { @@ -58,6 +27,19 @@ func (s *ServerCommon) OpenStreamTransport(ctx context.Context, transport *Trans if transport == nil { return nil, errStreamTransportNil } + return s.openStreamTransport(ctx, transport, opt) +} + +func (s *ServerCommon) openStreamTransport(ctx context.Context, transport *TransportConn, opt StreamOpenOptions) (Stream, error) { + if s == nil { + return nil, errStreamServerNil + } + if transport == nil { + return nil, errStreamTransportNil + } + if err := s.ensureServerTransportSendReady(transport); err != nil { + return nil, err + } logical := transport.LogicalConn() if logical == nil { return nil, errStreamLogicalConnNil @@ -71,26 +53,36 @@ func (s *ServerCommon) OpenStreamTransport(ctx context.Context, transport *Trans if _, exists := runtime.lookup(scope, req.StreamID); exists { return nil, errStreamAlreadyExists } - resp, err := sendStreamOpenServerTransport(ctx, s, transport, req) + dataID, err := runtime.reserveDataID(scope) if err != nil { return nil, err } - if resp.DataID != 0 { - req.DataID = resp.DataID + req.DataID = dataID + stream := newStreamHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, transport.TransportGeneration(), serverStreamCloseSender(s, logical, transport), serverStreamResetSender(s, logical, transport), serverStreamDataSender(s, transport), runtime.configSnapshot()) + if err := runtime.adoptReserved(scope, stream); err != nil { + runtime.releaseDataID(scope, req.DataID) + return nil, err + } + resp, err := sendStreamOpenServerTransport(ctx, s, transport, req) + if err != nil { + s.bestEffortStreamResetTransport(transport, StreamResetRequest{StreamID: req.StreamID, DataID: req.DataID, Error: err.Error()}) + stream.markReset(err) + return nil, err + } + if resp.DataID != 0 && resp.DataID != req.DataID { + err = errStreamAlreadyExists + s.bestEffortStreamResetTransport(transport, StreamResetRequest{StreamID: req.StreamID, Error: "stream data id mismatch"}) + stream.markReset(err) + return nil, err } if resp.FastPathVersion != 0 { - req.FastPathVersion = resp.FastPathVersion + stream.setFastPathVersion(resp.FastPathVersion) } else { - req.FastPathVersion = streamFastPathVersionV1 + stream.setFastPathVersion(streamFastPathVersionV1) } - req.Metadata = mergeStreamMetadata(req.Metadata, resp.Metadata) - stream := newStreamHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, resp.TransportGeneration, serverStreamCloseSender(s, logical, transport), serverStreamResetSender(s, logical, transport), serverStreamDataSender(s, transport), runtime.configSnapshot()) - if err := runtime.register(scope, stream); err != nil { - _, _ = sendStreamResetServerTransport(context.Background(), s, transport, StreamResetRequest{ - StreamID: req.StreamID, - Error: err.Error(), - }) - return nil, err + stream.metadata = mergeStreamMetadata(req.Metadata, resp.Metadata) + if resp.TransportGeneration != 0 { + stream.setTransportGeneration(resp.TransportGeneration) } return stream, nil } @@ -114,13 +106,14 @@ func serverStreamCloseSender(s *ServerCommon, logical *LogicalConn, transport *T return func(ctx context.Context, stream *streamHandle, full bool) error { req := StreamCloseRequest{ StreamID: stream.ID(), + DataID: stream.dataIDSnapshot(), Full: full, } - if logical != nil { - _, err := sendStreamCloseServerLogical(ctx, s, logical, req) + if transport != nil { + _, err := sendStreamCloseServerTransport(ctx, s, transport, req) return err } - _, err := sendStreamCloseServerTransport(ctx, s, transport, req) + _, err := sendStreamCloseServerLogical(ctx, s, logical, req) return err } } @@ -128,18 +121,29 @@ func serverStreamCloseSender(s *ServerCommon, logical *LogicalConn, transport *T func serverStreamResetSender(s *ServerCommon, logical *LogicalConn, transport *TransportConn) streamResetSender { return func(ctx context.Context, stream *streamHandle, message string) error { req := StreamResetRequest{ - StreamID: stream.ID(), - Error: message, + StreamID: stream.ID(), + DataID: stream.dataIDSnapshot(), + Error: message, + RecordFailure: stream.recordResetFailure(), } - if logical != nil { - _, err := sendStreamResetServerLogical(ctx, s, logical, req) + if transport != nil { + _, err := sendStreamResetServerTransport(ctx, s, transport, req) return err } - _, err := sendStreamResetServerTransport(ctx, s, transport, req) + _, err := sendStreamResetServerLogical(ctx, s, logical, req) return err } } +func (s *ServerCommon) bestEffortStreamResetTransport(transport *TransportConn, req StreamResetRequest) { + if s == nil || transport == nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), streamDispatchRejectTimeout) + defer cancel() + _, _ = sendStreamResetServerTransport(ctx, s, transport, req) +} + func serverStreamDataSender(s *ServerCommon, transport *TransportConn) streamDataSender { return func(ctx context.Context, stream *streamHandle, chunk []byte) error { if s == nil { diff --git a/session_owner_state.go b/session_owner_state.go index 327e3ff..66e24bf 100644 --- a/session_owner_state.go +++ b/session_owner_state.go @@ -7,6 +7,7 @@ const ( ownerSessionStateStarting ownerSessionStateRunning ownerSessionStateStopping + ownerSessionStateFinalizing ownerSessionStateStopped ) @@ -43,7 +44,7 @@ func markOwnerSessionStarted(state *atomic.Int32) { switch current { case ownerSessionStateRunning: return - case ownerSessionStateStopping: + case ownerSessionStateStopping, ownerSessionStateFinalizing: return case ownerSessionStateStarting, ownerSessionStateIdle, ownerSessionStateStopped: if state.CompareAndSwap(current, ownerSessionStateRunning) { @@ -55,25 +56,47 @@ func markOwnerSessionStarted(state *atomic.Int32) { } } -func markOwnerSessionStopping(state *atomic.Int32) { +func markOwnerSessionStopping(state *atomic.Int32) bool { if state == nil { - return + return false } for { current := state.Load() switch current { - case ownerSessionStateStopping, ownerSessionStateStopped: - return + case ownerSessionStateStopping, ownerSessionStateFinalizing, ownerSessionStateStopped: + return false case ownerSessionStateRunning, ownerSessionStateStarting: if state.CompareAndSwap(current, ownerSessionStateStopping) { - return + return true } case ownerSessionStateIdle: if state.CompareAndSwap(current, ownerSessionStateStopped) { - return + return true } default: - return + return false + } + } +} + +// claimOwnerSessionStop elects exactly one caller to run session cleanup. +// finalizing remains externally visible as "stopping" while preventing a +// concurrent Stop/read-error path from entering cleanup a second time. +func claimOwnerSessionStop(state *atomic.Int32) bool { + if state == nil { + return false + } + for { + current := state.Load() + switch current { + case ownerSessionStateFinalizing, ownerSessionStateStopped: + return false + case ownerSessionStateIdle, ownerSessionStateStarting, ownerSessionStateRunning, ownerSessionStateStopping: + if state.CompareAndSwap(current, ownerSessionStateFinalizing) { + return true + } + default: + return false } } } @@ -101,7 +124,7 @@ func ownerSessionStateName(state int32) string { return "starting" case ownerSessionStateRunning: return "running" - case ownerSessionStateStopping: + case ownerSessionStateStopping, ownerSessionStateFinalizing: return "stopping" case ownerSessionStateStopped: return "stopped" @@ -145,11 +168,18 @@ func (c *ClientCommon) markClientSessionStopped() { markOwnerSessionStopped(&c.sessionOwnerState) } -func (c *ClientCommon) markClientSessionStopping() { +func (c *ClientCommon) markClientSessionStopping() bool { if c == nil { - return + return false } - markOwnerSessionStopping(&c.sessionOwnerState) + return markOwnerSessionStopping(&c.sessionOwnerState) +} + +func (c *ClientCommon) claimClientSessionStop() bool { + if c == nil { + return false + } + return claimOwnerSessionStop(&c.sessionOwnerState) } func (c *ClientCommon) ownerSessionState() int32 { @@ -191,11 +221,18 @@ func (s *ServerCommon) markServerSessionStopped() { markOwnerSessionStopped(&s.sessionOwnerState) } -func (s *ServerCommon) markServerSessionStopping() { +func (s *ServerCommon) markServerSessionStopping() bool { if s == nil { - return + return false } - markOwnerSessionStopping(&s.sessionOwnerState) + return markOwnerSessionStopping(&s.sessionOwnerState) +} + +func (s *ServerCommon) claimServerSessionStop() bool { + if s == nil { + return false + } + return claimOwnerSessionStop(&s.sessionOwnerState) } func (s *ServerCommon) ownerSessionState() int32 { diff --git a/session_owner_state_test.go b/session_owner_state_test.go index 826b39c..a581864 100644 --- a/session_owner_state_test.go +++ b/session_owner_state_test.go @@ -1,6 +1,39 @@ package notify -import "testing" +import ( + "sync" + "sync/atomic" + "testing" +) + +func TestClaimOwnerSessionStopHasSingleCleanupWinner(t *testing.T) { + var state atomic.Int32 + state.Store(ownerSessionStateRunning) + + const contenders = 32 + start := make(chan struct{}) + var wg sync.WaitGroup + var winners atomic.Int32 + wg.Add(contenders) + for range contenders { + go func() { + defer wg.Done() + <-start + if claimOwnerSessionStop(&state) { + winners.Add(1) + } + }() + } + close(start) + wg.Wait() + + if got := winners.Load(); got != 1 { + t.Fatalf("session cleanup winners = %d, want 1", got) + } + if got := ownerSessionStateName(state.Load()); got != "stopping" { + t.Fatalf("claimed session state = %q, want stopping", got) + } +} func TestClientOwnerSessionStateStartRollback(t *testing.T) { client := NewClient().(*ClientCommon) diff --git a/session_state.go b/session_state.go index 3a1c892..8770b9c 100644 --- a/session_state.go +++ b/session_state.go @@ -92,7 +92,9 @@ func (c *ClientCommon) markSessionStarted() { } func (c *ClientCommon) markSessionStopped(reason string, err error) { - c.markClientSessionStopping() + if !c.claimClientSessionStop() { + return + } sessionMarkStopped(&c.alive, &c.mu, &c.status, reason, err, c.clientStopFuncSnapshot(), c.clearClientSessionRuntimeTransport, c.clearClientSessionRuntimeQueue, @@ -107,7 +109,9 @@ func (s *ServerCommon) markSessionStarted() { } func (s *ServerCommon) markSessionStopped(reason string, err error) { - s.markServerSessionStopping() + if !s.claimServerSessionStop() { + return + } sessionMarkStopped(&s.alive, &s.mu, &s.status, reason, err, s.serverStopFuncSnapshot(), s.clearServerSessionRuntimeTransport, s.clearServerSessionRuntimeQueue, diff --git a/signal_reliable.go b/signal_reliable.go index 0b07f7d..02cbc3f 100644 --- a/signal_reliable.go +++ b/signal_reliable.go @@ -230,11 +230,15 @@ func sendSignalWithAckTracked(state *signalReliabilityState, scope string, signa } func (c *ClientCommon) sendSignalEnvelopeMaybeReliable(env Envelope, msg TransferMsg) error { + return c.sendSignalEnvelopeMaybeReliableAtRoute(c.clientSessionRouteSnapshot(), env, msg) +} + +func (c *ClientCommon) sendSignalEnvelopeMaybeReliableAtRoute(route clientSessionRoute, env Envelope, msg TransferMsg) error { state := c.getSignalReliabilityState() state.incSignalSend() cfg := c.getSignalReliabilityConfig() if !cfg.Enabled || !signalCanUseTransportAck(msg) { - return c.sendEnvelope(env) + return c.sendEnvelopeAtRoute(route, env) } state.incReliableSend() return retryReliableSignalSendWithAttempt(cfg, func(cfg signalReliabilityConfig, attempt int) error { @@ -242,7 +246,7 @@ func (c *ClientCommon) sendSignalEnvelopeMaybeReliable(env Envelope, msg Transfe state.incRetry() } return sendSignalWithAckTracked(state, clientFileScope(), env.ID, cfg.AckTimeout, c.getSignalAckPool(), func() error { - return c.sendEnvelope(env) + return c.sendEnvelopeAtRoute(route, env) }) }) } @@ -274,7 +278,11 @@ func (s *ServerCommon) sendSignalEnvelopeMaybeReliableTransport(transport *Trans } func (c *ClientCommon) sendSignalAck(signalID uint64) error { - return c.sendEnvelope(newSignalAckEnvelope(signalID)) + return c.sendSignalAckAtRoute(c.clientSessionRouteSnapshot(), signalID) +} + +func (c *ClientCommon) sendSignalAckAtRoute(route clientSessionRoute, signalID uint64) error { + return c.sendEnvelopeAtRoute(route, newSignalAckEnvelope(signalID)) } func (s *ServerCommon) sendSignalAck(logical *LogicalConn, signalID uint64) error { @@ -311,6 +319,10 @@ func (s *ServerCommon) handleSignalAckEnvelopeTransport(transport *TransportConn } func (c *ClientCommon) handleReceivedSignalReliability(msg TransferMsg) bool { + return c.handleReceivedSignalReliabilityAtRoute(c.clientSessionRouteSnapshot(), msg) +} + +func (c *ClientCommon) handleReceivedSignalReliabilityAtRoute(route clientSessionRoute, msg TransferMsg) bool { cfg := c.getSignalReliabilityConfig() if !cfg.Enabled || !signalCanUseTransportAck(msg) { return false @@ -321,7 +333,7 @@ func (c *ClientCommon) handleReceivedSignalReliability(msg TransferMsg) bool { state.incDuplicateRecv() } state.incAckSend() - if err := c.sendSignalAck(msg.ID); err != nil { + if err := c.sendSignalAckAtRoute(route, msg.ID); err != nil { state.incAckSendError() if c.showError || c.debugMode { fmt.Println("client send signal ack error", err) diff --git a/stream.go b/stream.go index a9b966b..99a9752 100644 --- a/stream.go +++ b/stream.go @@ -80,6 +80,7 @@ var ( errStreamNotFound = errors.New("stream not found") errStreamHandlerNotConfigured = errors.New("stream handler is not configured") errStreamDataPathNotReady = errors.New("stream data path is not implemented yet") + errStreamDataIDExhausted = errors.New("stream data id exhausted") errStreamRejected = errors.New("stream open rejected") errStreamReset = errors.New("stream reset") errStreamBackpressureExceeded = errors.New("stream inbound backpressure exceeded") @@ -153,6 +154,7 @@ type streamHandle struct { channel StreamChannel metadata StreamMetadata sessionEpoch uint64 + clientRoute clientSessionRoute client *ClientCommon logical *LogicalConn transport *TransportConn @@ -173,6 +175,9 @@ type streamHandle struct { writeMu sync.Mutex mu sync.Mutex + negotiationMu sync.RWMutex + finalizeOnce sync.Once + acceptState atomic.Uint32 // 0=pending, 1=dispatched/handled, 2=reset before dispatch localClosed bool localReadClosed bool remoteClosed bool @@ -259,14 +264,43 @@ func (s *streamHandle) acceptsClientSessionEpoch(epoch uint64) bool { return s.sessionEpoch == epoch } +func (s *streamHandle) setClientSessionRoute(route clientSessionRoute) { + if s == nil { + return + } + s.sessionEpoch = route.epoch + s.clientRoute = route +} + +func (s *streamHandle) clientSessionRouteSnapshot() clientSessionRoute { + if s == nil { + return clientSessionRoute{} + } + if !s.clientRoute.bound() && s.client != nil { + return s.client.clientSessionRouteSnapshot() + } + return s.clientRoute +} + +func (s *streamHandle) acceptsClientSessionRoute(route clientSessionRoute) bool { + if !s.acceptsClientSessionEpoch(route.epoch) { + return false + } + if s == nil || s.clientRoute.binding == nil || route.binding == nil { + return true + } + return s.clientRoute.binding == route.binding +} + func (s *streamHandle) acceptsTransportGeneration(transport *TransportConn) bool { if s == nil { return false } - if s.transportGeneration == 0 || transport == nil { + generation := s.TransportGeneration() + if generation == 0 || transport == nil { return true } - return s.transportGeneration == transport.TransportGeneration() + return generation == transport.TransportGeneration() } func (s *streamHandle) ID() string { @@ -302,9 +336,20 @@ func (s *streamHandle) fastPathVersionSnapshot() uint8 { if s == nil { return streamFastPathVersionV1 } + s.negotiationMu.RLock() + defer s.negotiationMu.RUnlock() return normalizeStreamFastPathVersion(s.fastPathVersion) } +func (s *streamHandle) setFastPathVersion(version uint8) { + if s == nil { + return + } + s.negotiationMu.Lock() + s.fastPathVersion = normalizeStreamFastPathVersion(version) + s.negotiationMu.Unlock() +} + func (s *streamHandle) Channel() StreamChannel { if s == nil { return StreamDataChannel @@ -344,9 +389,57 @@ func (s *streamHandle) TransportGeneration() uint64 { if s == nil { return 0 } + s.negotiationMu.RLock() + defer s.negotiationMu.RUnlock() return s.transportGeneration } +func (s *streamHandle) setTransportGeneration(generation uint64) { + if s == nil || generation == 0 { + return + } + s.negotiationMu.Lock() + s.transportGeneration = generation + s.negotiationMu.Unlock() +} + +func (s *streamHandle) acceptsCurrentTransport() bool { + if s == nil { + return false + } + if s.transport != nil { + return s.transport.IsCurrent() + } + if s.client != nil { + return s.client.clientSessionRouteCurrent(s.clientSessionRouteSnapshot()) + } + return true +} + +func (s *streamHandle) acceptDispatchAllowed() bool { + if s == nil || s.acceptState.Load() == 2 { + return false + } + if err := s.resetErrSnapshot(); err != nil { + return false + } + return s.acceptsCurrentTransport() +} + +func (s *streamHandle) claimAcceptDispatch() bool { + if s == nil { + return false + } + return s.acceptState.CompareAndSwap(0, 1) +} + +func (s *streamHandle) markAcceptHandled() { + if s == nil { + return + } + s.acceptState.CompareAndSwap(0, 1) +} + func (s *streamHandle) LocalAddr() net.Addr { if s == nil { return nil @@ -780,20 +873,28 @@ func (s *streamHandle) Reset(err error) error { return err } resetFn := s.resetFn + deadline := s.effectiveWriteDeadlineLocked(time.Now(), s.writeTimeout) s.mu.Unlock() - - if resetFn != nil { - ctx, cancel, err := s.newControlContext() - if err != nil { - return err - } - defer cancel() - if sendErr := resetFn(ctx, s, streamResetMessage(resetErr)); sendErr != nil { - return sendErr - } + if !s.applyResetState(resetErr) { + return s.resetErrSnapshot() + } + if resetFn == nil { + return nil + } + limit := time.Now().Add(5 * time.Second) + if deadline.IsZero() || deadline.After(limit) { + deadline = limit + } + ctx, cancel := context.WithDeadline(context.Background(), deadline) + defer cancel() + done := make(chan error, 1) + go func() { done <- resetFn(ctx, s, streamResetMessage(resetErr)) }() + select { + case err := <-done: + return err + case <-ctx.Done(): + return ctx.Err() } - s.markReset(resetErr) - return nil } func (s *streamHandle) markRemoteClosed() { @@ -826,17 +927,25 @@ func (s *streamHandle) markPeerClosed() { } func (s *streamHandle) markReset(err error) { + _ = s.applyResetState(err) +} + +func (s *streamHandle) applyResetState(err error) bool { if s == nil { - return + return false } + s.acceptState.CompareAndSwap(0, 2) s.mu.Lock() - if s.resetErr == nil { - s.resetErr = streamResetError(err) - s.clearBufferedDataLocked() + if s.resetErr != nil { + s.mu.Unlock() + return false } + s.resetErr = streamResetError(err) + s.clearBufferedDataLocked() s.mu.Unlock() s.notifyReadable() s.finalize() + return true } func (s *streamHandle) resetErrSnapshot() error { @@ -1019,7 +1128,7 @@ func (s *streamHandle) snapshot() StreamSnapshot { Channel: s.channel, Metadata: cloneStreamMetadata(s.metadata), SessionEpoch: s.sessionEpoch, - TransportGeneration: s.transportGeneration, + TransportGeneration: s.TransportGeneration(), LocalClosed: s.localClosed, LocalReadClosed: s.localReadClosed, RemoteClosed: s.remoteClosed, @@ -1059,7 +1168,7 @@ func (s *streamHandle) snapshot() StreamSnapshot { var diag snapshotBindingDiagnostics switch { case s.logical != nil || s.transport != nil: - diag = snapshotBindingDiagnosticsFromLogical(s.logical, s.transport, s.transportGeneration) + diag = snapshotBindingDiagnosticsFromLogical(s.logical, s.transport, s.TransportGeneration()) case s.client != nil: diag = snapshotBindingDiagnosticsFromClient(s.client, s.sessionEpoch) } @@ -1095,12 +1204,14 @@ func (s *streamHandle) finalize() { if s == nil { return } - if s.cancel != nil { - s.cancel() - } - if s.runtime != nil { - s.runtime.remove(s.runtimeScope, s.id) - } + s.finalizeOnce.Do(func() { + if s.cancel != nil { + s.cancel() + } + if s.runtime != nil { + s.runtime.remove(s.runtimeScope, s) + } + }) } func (s *streamHandle) waitReadable(ctx context.Context, notify <-chan struct{}, deadlineNotify <-chan struct{}, deadline time.Time) error { diff --git a/stream_batch_sender.go b/stream_batch_sender.go index 13a5a3f..3340e98 100644 --- a/stream_batch_sender.go +++ b/stream_batch_sender.go @@ -46,11 +46,14 @@ type streamBatchSender struct { stopCh chan struct{} doneCh chan struct{} - stopOnce sync.Once - flushMu sync.Mutex - queued atomic.Int64 - errMu sync.Mutex - err error + stopOnce sync.Once + admissionMu sync.Mutex + admitting sync.WaitGroup + admissionClosed bool + flushMu sync.Mutex + queued atomic.Int64 + errMu sync.Mutex + err error } func newStreamBatchSender(binding *transportBinding, codec streamBatchCodec, writeTimeoutProvider func() time.Duration) *streamBatchSender { @@ -146,14 +149,12 @@ func (s *streamBatchSender) submitRequest(req streamBatchRequest) error { } } s.queued.Add(1) - select { - case <-req.ctx.Done(): - s.queued.Add(-1) - return normalizeStreamDeadlineError(req.ctx.Err()) - case <-s.stopCh: + if !s.enqueue(req) { s.queued.Add(-1) + if err := req.ctx.Err(); err != nil { + return normalizeStreamDeadlineError(err) + } return s.stoppedErr() - case s.reqCh <- req: } select { case err := <-req.done: @@ -210,7 +211,11 @@ func (s *streamBatchSender) tryDirectSubmit(req streamBatchRequest) (bool, error return true, err } if err := s.flush([]streamBatchRequest{req}); err != nil { - s.setErr(err) + if isBatchSenderQueueWaitError(err) { + return true, err + } + s.markFailed(err) + s.waitAdmissions() s.failPending(err) return true, err } @@ -240,7 +245,10 @@ func (s *streamBatchSender) run() { if timerCh == nil { select { case <-s.stopCh: - s.failPending(s.stoppedErr()) + err := s.stoppedErr() + s.waitAdmissions() + s.failBatch(batch, err) + s.failPending(err) return case next := <-s.reqCh: batch = append(batch, next) @@ -255,7 +263,10 @@ func (s *streamBatchSender) run() { if timer != nil { timer.Stop() } - s.failPending(s.stoppedErr()) + err := s.stoppedErr() + s.waitAdmissions() + s.failBatch(batch, err) + s.failPending(err) return case next := <-s.reqCh: batch = append(batch, next) @@ -296,10 +307,17 @@ func (s *streamBatchSender) run() { } s.flushMu.Unlock() if err != nil { - s.setErr(err) + if isBatchSenderQueueWaitError(err) { + for _, item := range active { + s.finishRequest(item, err) + } + continue + } + s.markFailed(err) for _, item := range active { s.finishRequest(item, err) } + s.waitAdmissions() s.failPending(err) return } @@ -312,6 +330,7 @@ func (s *streamBatchSender) run() { func (s *streamBatchSender) nextRequest() (streamBatchRequest, bool) { select { case <-s.stopCh: + s.waitAdmissions() s.failPending(s.stoppedErr()) return streamBatchRequest{}, false case req := <-s.reqCh: @@ -347,6 +366,9 @@ func (s *streamBatchSender) flush(requests []streamBatchRequest) error { lockAcquired, err := s.binding.withConnWriteLockContextStopDeadlineManaged(writeCtx, s.stopCh, writeDeadlineFromTimeout(writeTimeout), func(conn net.Conn) error { return writeFramedPayloadBatchUnlocked(conn, queue, payloads) }) + if !lockAcquired && isBatchSenderQueueWaitCause(err) { + return newBatchSenderQueueWaitError(err) + } s.binding.observeStreamAdaptivePayloadWrite(payloadBytes, time.Since(started), writeTimeout, err) if lockAcquired && err != nil { // A failed framed write may have emitted only part of a frame. @@ -522,10 +544,8 @@ func (s *streamBatchSender) stop() { if s == nil { return } - s.stopOnce.Do(func() { - s.setErr(errTransportDetached) - close(s.stopCh) - }) + s.markFailed(errTransportDetached) + s.waitAdmissions() <-s.doneCh // Direct submissions flush on the caller goroutine rather than run(). Wait // for that path too before declaring the binding safe to hand off. @@ -533,6 +553,28 @@ func (s *streamBatchSender) stop() { s.flushMu.Unlock() } +func (s *streamBatchSender) enqueue(req streamBatchRequest) bool { + if s == nil { + return false + } + s.admissionMu.Lock() + if s.admissionClosed { + s.admissionMu.Unlock() + return false + } + s.admitting.Add(1) + s.admissionMu.Unlock() + defer s.admitting.Done() + select { + case <-req.ctx.Done(): + return false + case <-s.stopCh: + return false + case s.reqCh <- req: + return true + } +} + func (s *streamBatchSender) failPending(err error) { for { select { @@ -544,6 +586,12 @@ func (s *streamBatchSender) failPending(err error) { } } +func (s *streamBatchSender) failBatch(batch []streamBatchRequest, err error) { + for _, item := range batch { + s.finishRequest(item, err) + } +} + func (s *streamBatchSender) setErr(err error) { if s == nil || err == nil { return @@ -555,6 +603,25 @@ func (s *streamBatchSender) setErr(err error) { s.errMu.Unlock() } +func (s *streamBatchSender) markFailed(err error) { + if s == nil { + return + } + s.setErr(err) + s.stopOnce.Do(func() { + s.admissionMu.Lock() + s.admissionClosed = true + close(s.stopCh) + s.admissionMu.Unlock() + }) +} + +func (s *streamBatchSender) waitAdmissions() { + if s != nil { + s.admitting.Wait() + } +} + func (s *streamBatchSender) errSnapshot() error { if s == nil { return errTransportDetached diff --git a/stream_control.go b/stream_control.go index 677536d..47ea4b6 100644 --- a/stream_control.go +++ b/stream_control.go @@ -3,6 +3,7 @@ package notify import ( "context" "errors" + "strings" "time" ) @@ -28,6 +29,7 @@ type StreamOpenResponse struct { type StreamCloseRequest struct { StreamID string + DataID uint64 Full bool } @@ -38,9 +40,10 @@ type StreamCloseResponse struct { } type StreamResetRequest struct { - StreamID string - DataID uint64 - Error string + StreamID string + DataID uint64 + Error string + RecordFailure *RecordFailure } type StreamResetResponse struct { @@ -87,6 +90,15 @@ func (c *ClientCommon) handleInboundStreamOpen(msg *Message) { replyStreamControlIfNeeded(msg, resp) return } + route := msg.clientRoute + if !route.bound() { + route = c.clientSessionRouteSnapshot() + } + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + resp.Error = err.Error() + replyStreamControlIfNeeded(msg, resp) + return + } runtime := c.getStreamRuntime() if runtime == nil { resp.Error = errStreamRuntimeNil.Error() @@ -94,17 +106,38 @@ func (c *ClientCommon) handleInboundStreamOpen(msg *Message) { return } scope := clientFileScope() + if existing, ok := runtime.lookup(scope, req.StreamID); ok && !existing.acceptsClientSessionRoute(route) { + existing.markReset(transportDetachedSessionEpochError()) + } req.FastPathVersion = negotiateStreamFastPathVersion(req.FastPathVersion) resp.FastPathVersion = req.FastPathVersion - if req.DataID == 0 { - req.DataID = runtime.nextDataID() - resp.DataID = req.DataID - } req.Metadata, resp.Metadata = negotiateRecordStreamOpenMetadata(req.Channel, req.Metadata) - stream := newStreamHandle(c.clientStopContextSnapshot(), runtime, scope, req, c.currentClientSessionEpoch(), nil, nil, 0, clientStreamCloseSender(c), clientStreamResetSender(c), clientStreamDataSender(c, c.currentClientSessionEpoch()), runtime.configSnapshot()) + parent := clientSessionRouteContext(route) + if parent == nil { + parent = c.clientStopContextSnapshot() + } + stream := newStreamHandle(parent, runtime, scope, req, route.epoch, nil, nil, 0, clientStreamCloseSender(c), clientStreamResetSender(c), clientStreamDataSender(c, route), runtime.configSnapshot()) stream.setClientSnapshotOwner(c) - stream.setAddrSnapshot(c.clientStreamAddrSnapshot()) - if err := runtime.register(scope, stream); err != nil { + stream.setClientSessionRoute(route) + stream.setAddrSnapshot(c.clientStreamAddrSnapshotAtRoute(route)) + if err := runtime.adoptInbound(scope, stream); err != nil { + resp.Error = err.Error() + replyStreamControlIfNeeded(msg, resp) + return + } + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + runtime.remove(scope, stream) + stream.markReset(err) + resp.Error = err.Error() + replyStreamControlIfNeeded(msg, resp) + return + } + if !stream.acceptDispatchAllowed() || !stream.claimAcceptDispatch() { + err := stream.resetErrSnapshot() + if err == nil { + err = transportDetachedSessionEpochError() + stream.markReset(err) + } resp.Error = err.Error() replyStreamControlIfNeeded(msg, resp) return @@ -117,6 +150,7 @@ func (c *ClientCommon) handleInboundStreamOpen(msg *Message) { return } resp.Accepted = true + resp.DataID = stream.dataIDSnapshot() resp.TransportGeneration = stream.TransportGeneration() replyStreamControlIfNeeded(msg, resp) return @@ -129,6 +163,7 @@ func (c *ClientCommon) handleInboundStreamOpen(msg *Message) { return } resp.Accepted = true + resp.DataID = stream.dataIDSnapshot() resp.TransportGeneration = stream.TransportGeneration() replyStreamControlIfNeeded(msg, resp) return @@ -181,16 +216,37 @@ func (s *ServerCommon) handleInboundStreamOpen(msg *Message) { return } transport := messageTransportConnSnapshot(msg) + if transport != nil && !transport.IsCurrent() { + resp.Error = transportDetachedErrorForTransport(transport).Error() + replyStreamControlIfNeeded(msg, resp) + return + } scope := serverFileScope(logical) + if existing, ok := runtime.lookup(scope, req.StreamID); ok && !existing.acceptsTransportGeneration(transport) { + existing.markReset(transportDetachedGenerationMismatchError(existing.TransportGeneration(), transport)) + } req.FastPathVersion = negotiateStreamFastPathVersion(req.FastPathVersion) resp.FastPathVersion = req.FastPathVersion - if req.DataID == 0 { - req.DataID = runtime.nextDataID() - resp.DataID = req.DataID - } req.Metadata, resp.Metadata = negotiateRecordStreamOpenMetadata(req.Channel, req.Metadata) stream := newStreamHandle(logical.stopContextSnapshot(), runtime, scope, req, 0, logical, transport, streamTransportGeneration(logical, transport), serverStreamCloseSender(s, logical, transport), serverStreamResetSender(s, logical, transport), serverStreamDataSender(s, transport), runtime.configSnapshot()) - if err := runtime.register(scope, stream); err != nil { + if err := runtime.adoptInbound(scope, stream); err != nil { + resp.Error = err.Error() + replyStreamControlIfNeeded(msg, resp) + return + } + if transport != nil && !transport.IsCurrent() { + runtime.remove(scope, stream) + stream.markReset(transportDetachedErrorForTransport(transport)) + resp.Error = transportDetachedErrorForTransport(transport).Error() + replyStreamControlIfNeeded(msg, resp) + return + } + if !stream.acceptDispatchAllowed() || !stream.claimAcceptDispatch() { + err := stream.resetErrSnapshot() + if err == nil { + err = transportDetachedErrorForTransport(transport) + stream.markReset(err) + } resp.Error = err.Error() replyStreamControlIfNeeded(msg, resp) return @@ -203,6 +259,7 @@ func (s *ServerCommon) handleInboundStreamOpen(msg *Message) { return } resp.Accepted = true + resp.DataID = stream.dataIDSnapshot() resp.TransportGeneration = stream.TransportGeneration() replyStreamControlIfNeeded(msg, resp) return @@ -215,6 +272,7 @@ func (s *ServerCommon) handleInboundStreamOpen(msg *Message) { return } resp.Accepted = true + resp.DataID = stream.dataIDSnapshot() resp.TransportGeneration = stream.TransportGeneration() replyStreamControlIfNeeded(msg, resp) return @@ -262,12 +320,21 @@ func (c *ClientCommon) handleInboundStreamClose(msg *Message) { replyStreamControlIfNeeded(msg, resp) return } - stream, ok := runtime.lookup(clientFileScope(), req.StreamID) + route := msg.clientRoute + if !route.bound() { + route = c.clientSessionRouteSnapshot() + } + stream, ok := runtime.lookupControl(clientFileScope(), req.StreamID, req.DataID) if !ok { resp.Error = errStreamNotFound.Error() replyStreamControlIfNeeded(msg, resp) return } + if !stream.acceptsClientSessionRoute(route) { + resp.Error = transportDetachedSessionEpochError().Error() + replyStreamControlIfNeeded(msg, resp) + return + } if req.Full { stream.markPeerClosed() } else { @@ -293,12 +360,17 @@ func (s *ServerCommon) handleInboundStreamClose(msg *Message) { } logical := messageLogicalConnSnapshot(msg) scope := serverFileScope(logical) - stream, ok := runtime.lookup(scope, req.StreamID) + stream, ok := runtime.lookupControl(scope, req.StreamID, req.DataID) if !ok { resp.Error = errStreamNotFound.Error() replyStreamControlIfNeeded(msg, resp) return } + if !stream.acceptsTransportGeneration(messageTransportConnSnapshot(msg)) { + resp.Error = transportDetachedGenerationMismatchError(stream.TransportGeneration(), messageTransportConnSnapshot(msg)).Error() + replyStreamControlIfNeeded(msg, resp) + return + } if req.Full { stream.markPeerClosed() } else { @@ -322,19 +394,25 @@ func (c *ClientCommon) handleInboundStreamReset(msg *Message) { replyStreamControlIfNeeded(msg, resp) return } - stream, ok := runtime.lookup(clientFileScope(), req.StreamID) - if !ok && req.DataID != 0 { - stream, ok = runtime.lookupByDataID(clientFileScope(), req.DataID) + route := msg.clientRoute + if !route.bound() { + route = c.clientSessionRouteSnapshot() } + stream, ok := runtime.lookupControl(clientFileScope(), req.StreamID, req.DataID) if !ok { resp.Error = errStreamNotFound.Error() replyStreamControlIfNeeded(msg, resp) return } + if !stream.acceptsClientSessionRoute(route) { + resp.Error = transportDetachedSessionEpochError().Error() + replyStreamControlIfNeeded(msg, resp) + return + } if resp.StreamID == "" { resp.StreamID = stream.ID() } - stream.markReset(streamResetError(streamRemoteResetError(req.Error))) + stream.markReset(req.resetError(stream.Channel())) resp.Accepted = true replyStreamControlIfNeeded(msg, resp) } @@ -355,19 +433,21 @@ func (s *ServerCommon) handleInboundStreamReset(msg *Message) { } logical := messageLogicalConnSnapshot(msg) scope := serverFileScope(logical) - stream, ok := runtime.lookup(scope, req.StreamID) - if !ok && req.DataID != 0 { - stream, ok = runtime.lookupByDataID(scope, req.DataID) - } + stream, ok := runtime.lookupControl(scope, req.StreamID, req.DataID) if !ok { resp.Error = errStreamNotFound.Error() replyStreamControlIfNeeded(msg, resp) return } + if !stream.acceptsTransportGeneration(messageTransportConnSnapshot(msg)) { + resp.Error = transportDetachedGenerationMismatchError(stream.TransportGeneration(), messageTransportConnSnapshot(msg)).Error() + replyStreamControlIfNeeded(msg, resp) + return + } if resp.StreamID == "" { resp.StreamID = stream.ID() } - stream.markReset(streamResetError(streamRemoteResetError(req.Error))) + stream.markReset(req.resetError(stream.Channel())) resp.Accepted = true replyStreamControlIfNeeded(msg, resp) } @@ -390,6 +470,17 @@ func sendStreamOpenClient(ctx context.Context, c Client, req StreamOpenRequest) return decodeStreamOpenResponse(msg) } +func sendStreamOpenClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req StreamOpenRequest) (StreamOpenResponse, error) { + if c == nil { + return StreamOpenResponse{}, errStreamClientNil + } + msg, err := c.sendObjCtxAtRoute(ctx, route, StreamOpenSignalKey, req) + if err != nil { + return StreamOpenResponse{}, err + } + return decodeStreamOpenResponse(msg) +} + func sendStreamOpenServerLogical(ctx context.Context, s Server, logical *LogicalConn, req StreamOpenRequest) (StreamOpenResponse, error) { if s == nil { return StreamOpenResponse{}, errStreamServerNil @@ -429,6 +520,17 @@ func sendStreamCloseClient(ctx context.Context, c Client, req StreamCloseRequest return decodeStreamCloseResponse(msg) } +func sendStreamCloseClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req StreamCloseRequest) (StreamCloseResponse, error) { + if c == nil { + return StreamCloseResponse{}, errStreamClientNil + } + msg, err := c.sendObjCtxAtRoute(ctx, route, StreamCloseSignalKey, req) + if err != nil { + return StreamCloseResponse{}, err + } + return decodeStreamCloseResponse(msg) +} + func sendStreamCloseServerLogical(ctx context.Context, s Server, logical *LogicalConn, req StreamCloseRequest) (StreamCloseResponse, error) { if s == nil { return StreamCloseResponse{}, errStreamServerNil @@ -468,6 +570,17 @@ func sendStreamResetClient(ctx context.Context, c Client, req StreamResetRequest return decodeStreamResetResponse(msg) } +func sendStreamResetClientAtRoute(ctx context.Context, c *ClientCommon, route clientSessionRoute, req StreamResetRequest) (StreamResetResponse, error) { + if c == nil { + return StreamResetResponse{}, errStreamClientNil + } + msg, err := c.sendObjCtxAtRoute(ctx, route, StreamResetSignalKey, req) + if err != nil { + return StreamResetResponse{}, err + } + return decodeStreamResetResponse(msg) +} + func sendStreamResetServerLogical(ctx context.Context, s Server, logical *LogicalConn, req StreamResetRequest) (StreamResetResponse, error) { if s == nil { return StreamResetResponse{}, errStreamServerNil @@ -533,7 +646,7 @@ func decodeStreamResetRequest(msg *Message) (StreamResetRequest, error) { if err := msg.Value.Orm(&req); err != nil { return StreamResetRequest{}, err } - if req.StreamID == "" { + if req.StreamID == "" && req.DataID == 0 { return StreamResetRequest{}, errStreamIDEmpty } return req, nil @@ -580,6 +693,9 @@ func streamControlResultError(op string, accepted bool, message string, callErr } func streamControlMessageError(message string) error { + if message == errTransportDetached.Error() || strings.HasPrefix(message, errTransportDetached.Error()+":") { + return errTransportDetached + } switch message { case errStreamNotFound.Error(): return errStreamNotFound diff --git a/stream_dispatcher.go b/stream_dispatcher.go index 22f9775..59f0160 100644 --- a/stream_dispatcher.go +++ b/stream_dispatcher.go @@ -12,10 +12,21 @@ import ( const streamDispatchRejectTimeout = 300 * time.Millisecond func (c *ClientCommon) dispatchStreamEnvelope(env Envelope) { + route := c.clientSessionRouteSnapshot() + if route.epoch == 0 { + route.epoch = c.currentClientSessionEpoch() + } + c.dispatchStreamEnvelopeAtRoute(route, env) +} + +func (c *ClientCommon) dispatchStreamEnvelopeAtRoute(route clientSessionRoute, env Envelope) { streamID := env.Stream.StreamID if streamID == "" { return } + if route.binding != nil && !c.clientSessionRouteCurrent(route) { + return + } runtime := c.getStreamRuntime() if runtime == nil { return @@ -25,16 +36,16 @@ func (c *ClientCommon) dispatchStreamEnvelope(env Envelope) { if c.showError || c.debugMode { fmt.Println("client stream data for unknown stream", streamID) } - c.bestEffortRejectInboundStreamData(streamID, 0, errStreamNotFound.Error()) + c.bestEffortRejectInboundStreamDataAtRoute(route, streamID, 0, errStreamNotFound.Error()) return } - if !stream.acceptsClientSessionEpoch(c.currentClientSessionEpoch()) { + if !stream.acceptsClientSessionRoute(route) { if c.showError || c.debugMode { fmt.Println("client stream data rejected by stale session epoch", streamID) } detachErr := transportDetachedSessionEpochError() stream.markReset(detachErr) - c.bestEffortRejectInboundStreamData(streamID, stream.dataIDSnapshot(), detachErr.Error()) + c.bestEffortRejectInboundStreamDataAtRoute(route, streamID, stream.dataIDSnapshot(), detachErr.Error()) return } if err := stream.pushChunk(env.Stream.Chunk); err != nil { @@ -42,7 +53,7 @@ func (c *ClientCommon) dispatchStreamEnvelope(env Envelope) { fmt.Println("client stream push chunk error", err) } if !errors.Is(err, io.EOF) { - c.bestEffortRejectInboundStreamData(streamID, stream.dataIDSnapshot(), err.Error()) + c.bestEffortRejectInboundStreamDataAtRoute(route, streamID, stream.dataIDSnapshot(), err.Error()) } } } @@ -83,13 +94,28 @@ func (s *ServerCommon) dispatchStreamEnvelope(logical *LogicalConn, transport *T } func (c *ClientCommon) dispatchFastStreamData(frame streamFastDataFrame) { - c.dispatchFastStreamDataWithOwner(frame, nil) + route := c.clientSessionRouteSnapshot() + if route.epoch == 0 { + route.epoch = c.currentClientSessionEpoch() + } + c.dispatchFastStreamDataWithOwnerAtRoute(route, frame, nil) } func (c *ClientCommon) dispatchFastStreamDataWithOwner(frame streamFastDataFrame, owner *streamReadPayloadOwner) { + route := c.clientSessionRouteSnapshot() + if route.epoch == 0 { + route.epoch = c.currentClientSessionEpoch() + } + c.dispatchFastStreamDataWithOwnerAtRoute(route, frame, owner) +} + +func (c *ClientCommon) dispatchFastStreamDataWithOwnerAtRoute(route clientSessionRoute, frame streamFastDataFrame, owner *streamReadPayloadOwner) { if frame.DataID == 0 { return } + if route.binding != nil && !c.clientSessionRouteCurrent(route) { + return + } runtime := c.getStreamRuntime() if runtime == nil { return @@ -99,16 +125,16 @@ func (c *ClientCommon) dispatchFastStreamDataWithOwner(frame streamFastDataFrame if c.showError || c.debugMode { fmt.Println("client stream data for unknown data id", frame.DataID) } - c.bestEffortRejectInboundStreamData("", frame.DataID, errStreamNotFound.Error()) + c.bestEffortRejectInboundStreamDataAtRoute(route, "", frame.DataID, errStreamNotFound.Error()) return } - if !stream.acceptsClientSessionEpoch(c.currentClientSessionEpoch()) { + if !stream.acceptsClientSessionRoute(route) { if c.showError || c.debugMode { fmt.Println("client stream data rejected by stale session epoch", frame.DataID) } detachErr := transportDetachedSessionEpochError() stream.markReset(detachErr) - c.bestEffortRejectInboundStreamData(stream.ID(), frame.DataID, detachErr.Error()) + c.bestEffortRejectInboundStreamDataAtRoute(route, stream.ID(), frame.DataID, detachErr.Error()) return } var err error @@ -122,7 +148,7 @@ func (c *ClientCommon) dispatchFastStreamDataWithOwner(frame streamFastDataFrame fmt.Println("client stream push chunk error", err) } if !errors.Is(err, io.EOF) { - c.bestEffortRejectInboundStreamData(stream.ID(), frame.DataID, err.Error()) + c.bestEffortRejectInboundStreamDataAtRoute(route, stream.ID(), frame.DataID, err.Error()) } } } @@ -172,12 +198,16 @@ func (s *ServerCommon) dispatchFastStreamDataWithOwner(logical *LogicalConn, tra } func (c *ClientCommon) bestEffortRejectInboundStreamData(streamID string, dataID uint64, message string) { + c.bestEffortRejectInboundStreamDataAtRoute(c.clientSessionRouteSnapshot(), streamID, dataID, message) +} + +func (c *ClientCommon) bestEffortRejectInboundStreamDataAtRoute(route clientSessionRoute, streamID string, dataID uint64, message string) { if c == nil || (streamID == "" && dataID == 0) { return } ctx, cancel := context.WithTimeout(context.Background(), streamDispatchRejectTimeout) defer cancel() - _, _ = sendStreamResetClient(ctx, c, StreamResetRequest{ + _, _ = sendStreamResetClientAtRoute(ctx, c, route, StreamResetRequest{ StreamID: streamID, DataID: dataID, Error: message, diff --git a/stream_fastpath.go b/stream_fastpath.go index 606a18d..46a36eb 100644 --- a/stream_fastpath.go +++ b/stream_fastpath.go @@ -81,6 +81,9 @@ func encodeStreamFastDataFrameHeader(dst []byte, dataID uint64, seq uint64, payl if dataID == 0 { return errStreamFastDataIDEmpty } + if payloadLen < 0 || uint64(payloadLen) > uint64(^uint32(0)) { + return errStreamFastPayloadInvalid + } if len(dst) < streamFastPayloadHeaderLen { return errStreamFastPayloadInvalid } @@ -139,8 +142,8 @@ func decodeStreamFastDataFrame(payload []byte) (streamFastDataFrame, bool, error if payload[4] != streamFastPayloadVersion || payload[5] != streamFastPayloadTypeData { return streamFastDataFrame{}, true, errStreamFastPayloadInvalid } - dataLen := int(binary.BigEndian.Uint32(payload[24:28])) - if dataLen < 0 || len(payload) != streamFastPayloadHeaderLen+dataLen { + wireDataLen := binary.BigEndian.Uint32(payload[24:28]) + if uint64(len(payload)-streamFastPayloadHeaderLen) != uint64(wireDataLen) { return streamFastDataFrame{}, true, errStreamFastPayloadInvalid } dataID := binary.BigEndian.Uint64(payload[8:16]) @@ -194,12 +197,23 @@ func (c *ClientCommon) encodeFastStreamBatchPayload(frames []streamFastDataFrame } func (c *ClientCommon) sendFastStreamData(ctx context.Context, stream *streamHandle, chunk []byte) error { + return c.sendFastStreamDataToBinding(ctx, c.clientTransportBindingSnapshot(), stream, chunk) +} + +func (c *ClientCommon) sendFastStreamDataAtRoute(ctx context.Context, route clientSessionRoute, stream *streamHandle, chunk []byte) error { + if err := c.ensureClientSessionRouteSendReady(route); err != nil { + return err + } + return c.sendFastStreamDataToBinding(ctx, route.binding, stream, chunk) +} + +func (c *ClientCommon) sendFastStreamDataToBinding(ctx context.Context, binding *transportBinding, stream *streamHandle, chunk []byte) error { if stream == nil { return io.ErrClosedPipe } dataID := stream.dataIDSnapshot() fastPathVersion := stream.fastPathVersionSnapshot() - if binding := c.clientTransportBindingSnapshot(); binding != nil && streamFastPathSupportsBatch(fastPathVersion) { + if binding != nil && streamFastPathSupportsBatch(fastPathVersion) { if sender := binding.clientStreamBatchSenderSnapshot(c); sender != nil { if maxPayload := streamAdaptiveFramePayloadLimit(binding); maxPayload > 0 && len(chunk) > maxPayload { startSeq := stream.reserveOutboundDataSeqs(streamFastSplitFrameCount(len(chunk), maxPayload)) @@ -221,7 +235,7 @@ func (c *ClientCommon) sendFastStreamData(ctx context.Context, stream *streamHan if err != nil { return err } - return c.writePayloadToTransport(payload) + return c.writePayloadToTransportBindingContextTimeout(ctx, binding, payload, 0) } func (s *ServerCommon) encodeFastStreamPayloadLogical(logical *LogicalConn, frame streamFastDataFrame) ([]byte, error) { @@ -278,7 +292,7 @@ func (s *ServerCommon) sendFastStreamDataTransport(ctx context.Context, logical } dataID := stream.dataIDSnapshot() fastPathVersion := stream.fastPathVersionSnapshot() - if binding := logical.transportBindingSnapshot(); binding != nil && binding.queueSnapshot() != nil && streamFastPathSupportsBatch(fastPathVersion) { + if binding := serverTransportBindingSnapshot(logical, transport); binding != nil && binding.queueSnapshot() != nil && streamFastPathSupportsBatch(fastPathVersion) { if sender := binding.serverStreamBatchSenderSnapshot(logical); sender != nil { if maxPayload := streamAdaptiveFramePayloadLimit(binding); maxPayload > 0 && len(chunk) > maxPayload { startSeq := stream.reserveOutboundDataSeqs(streamFastSplitFrameCount(len(chunk), maxPayload)) diff --git a/stream_lifecycle_test.go b/stream_lifecycle_test.go new file mode 100644 index 0000000..eb9698c --- /dev/null +++ b/stream_lifecycle_test.go @@ -0,0 +1,64 @@ +package notify + +import ( + "context" + "errors" + "io" + "testing" +) + +func TestStreamRuntimeAdoptFailureFinalizesCandidate(t *testing.T) { + runtime := newStreamRuntime("cstr") + scope := clientFileScope() + existing := newStreamHandle(context.Background(), runtime, scope, StreamOpenRequest{ + StreamID: "duplicate", + DataID: 1, + }, 0, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) + if err := runtime.register(scope, existing); err != nil { + t.Fatalf("register existing stream: %v", err) + } + defer existing.markReset(io.ErrClosedPipe) + + candidate := newStreamHandle(context.Background(), runtime, scope, StreamOpenRequest{ + StreamID: "duplicate", + DataID: 2, + }, 0, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) + if err := runtime.adopt(scope, candidate); !errors.Is(err, errStreamAlreadyExists) { + t.Fatalf("adopt error = %v, want %v", err, errStreamAlreadyExists) + } + if err := candidate.resetErrSnapshot(); !errors.Is(err, errStreamAlreadyExists) { + t.Fatalf("candidate reset error = %v, want %v", err, errStreamAlreadyExists) + } + select { + case <-candidate.Context().Done(): + default: + t.Fatal("failed stream adoption left candidate context active") + } +} + +func TestStreamRuntimeStaleFinalizeDoesNotRemoveReplacement(t *testing.T) { + runtime := newStreamRuntime("cstr") + scope := clientFileScope() + old := newStreamHandle(context.Background(), runtime, scope, StreamOpenRequest{ + StreamID: "reused", + DataID: 1, + }, 0, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) + if err := runtime.register(scope, old); err != nil { + t.Fatalf("register old stream: %v", err) + } + old.markReset(errors.New("old failed")) + + replacement := newStreamHandle(context.Background(), runtime, scope, StreamOpenRequest{ + StreamID: "reused", + DataID: 2, + }, 0, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) + if err := runtime.register(scope, replacement); err != nil { + t.Fatalf("register replacement stream: %v", err) + } + defer replacement.markReset(io.ErrClosedPipe) + + old.markReset(errors.New("late duplicate reset")) + if got, ok := runtime.lookup(scope, "reused"); !ok || got != replacement { + t.Fatalf("replacement stream after stale finalize = %p/%v, want %p/true", got, ok, replacement) + } +} diff --git a/stream_route_dataid_test.go b/stream_route_dataid_test.go new file mode 100644 index 0000000..93474ee --- /dev/null +++ b/stream_route_dataid_test.go @@ -0,0 +1,415 @@ +package notify + +import ( + "b612.me/stario" + "context" + "io" + "math" + "net" + "sync" + "sync/atomic" + "testing" + "time" +) + +func TestStreamRuntimeSeparatesPeerDataIDNamespaces(t *testing.T) { + clientRuntime := newStreamRuntime("cstrm") + serverRuntime := newStreamRuntime("sstrm") + clientID := clientRuntime.nextDataID() + serverID := serverRuntime.nextDataID() + if clientID == serverID || clientID%2 != 1 || serverID%2 != 0 { + t.Fatalf("client/server stream data ids = %d/%d, want disjoint odd/even namespaces", clientID, serverID) + } +} + +func TestStreamRuntimeKeepsLegacyZeroDataIDControlsCompatible(t *testing.T) { + tests := []struct { + name string + rolePrefix string + wantParity uint64 + }{ + {name: "client receives server stream", rolePrefix: "cstrm", wantParity: 0}, + {name: "server receives client stream", rolePrefix: "sstrm", wantParity: 1}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + runtime := newStreamRuntime(tt.rolePrefix) + scope := "legacy-zero" + stream := newStreamHandle(context.Background(), runtime, scope, StreamOpenRequest{ + StreamID: "legacy-stream", + DataID: 0, + }, 0, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) + if err := runtime.registerInbound(scope, stream); err != nil { + t.Fatalf("register legacy zero-DataID stream: %v", err) + } + defer stream.markReset(io.ErrClosedPipe) + if dataID := stream.dataIDSnapshot(); dataID == 0 || dataID%2 != tt.wantParity { + t.Fatalf("legacy assigned DataID = %d, want non-zero parity %d", dataID, tt.wantParity) + } + if got, ok := runtime.lookupControl(scope, stream.ID(), 0); !ok || got != stream { + t.Fatalf("legacy DataID=0 control lookup = %p/%v, want %p/true", got, ok, stream) + } + }) + } +} + +func TestStreamOpenConcurrentlyFromBothPeers(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKServer failed: %v", err) + } + serverAccepted := make(chan StreamAcceptInfo, 1) + server.SetStreamHandler(func(info StreamAcceptInfo) error { + serverAccepted <- info + return nil + }) + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatalf("server Listen failed: %v", err) + } + defer func() { _ = server.Stop() }() + + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKClient failed: %v", err) + } + clientAccepted := make(chan StreamAcceptInfo, 1) + client.SetStreamHandler(func(info StreamAcceptInfo) error { + clientAccepted <- info + return nil + }) + if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { + t.Fatalf("client Connect failed: %v", err) + } + defer func() { _ = client.Stop() }() + logical := waitForTransferControlLogicalConn(t, server, 2*time.Second) + + type openResult struct { + stream Stream + err error + } + clientResult := make(chan openResult, 1) + serverResult := make(chan openResult, 1) + start := make(chan struct{}) + var ready sync.WaitGroup + ready.Add(2) + go func() { + ready.Done() + <-start + stream, err := client.OpenStream(context.Background(), StreamOpenOptions{Channel: StreamDataChannel}) + clientResult <- openResult{stream: stream, err: err} + }() + go func() { + ready.Done() + <-start + stream, err := server.OpenStreamLogical(context.Background(), logical, StreamOpenOptions{Channel: StreamDataChannel}) + serverResult <- openResult{stream: stream, err: err} + }() + ready.Wait() + close(start) + + clientOpen := <-clientResult + serverOpen := <-serverResult + if clientOpen.err != nil || serverOpen.err != nil { + t.Fatalf("concurrent stream opens failed: client=%v server=%v", clientOpen.err, serverOpen.err) + } + clientInbound := waitAcceptedStream(t, serverAccepted, 2*time.Second) + serverInbound := waitAcceptedStream(t, clientAccepted, 2*time.Second) + clientID := clientOpen.stream.(*streamHandle).dataIDSnapshot() + serverID := serverOpen.stream.(*streamHandle).dataIDSnapshot() + if clientID%2 != 1 || serverID%2 != 0 { + t.Fatalf("local stream data ids = %d/%d, want odd/even", clientID, serverID) + } + if clientInbound.DataID != clientID || serverInbound.DataID != serverID { + t.Fatalf("stream data ids differ across peers: client=%d/%d server=%d/%d", clientID, clientInbound.DataID, serverID, serverInbound.DataID) + } + if _, err := clientOpen.stream.Write([]byte("client")); err != nil { + t.Fatalf("client stream write failed: %v", err) + } + readStreamExactly(t, clientInbound.Stream, "client", 2*time.Second) + if _, err := serverOpen.stream.Write([]byte("server")); err != nil { + t.Fatalf("server stream write failed: %v", err) + } + readStreamExactly(t, serverInbound.Stream, "server", 2*time.Second) + _ = clientOpen.stream.Close() + _ = clientInbound.Stream.Close() + _ = serverOpen.stream.Close() + _ = serverInbound.Stream.Close() +} + +func TestStreamHandlerWriteBeforeOpenReplyIsDelivered(t *testing.T) { + server := NewServer().(*ServerCommon) + if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKServer failed: %v", err) + } + server.SetStreamHandler(func(info StreamAcceptInfo) error { + _, err := info.Stream.Write([]byte("early")) + return err + }) + if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { + t.Fatalf("server Listen failed: %v", err) + } + defer func() { _ = server.Stop() }() + + client := NewClient().(*ClientCommon) + if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { + t.Fatalf("UseModernPSKClient failed: %v", err) + } + if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { + t.Fatalf("client Connect failed: %v", err) + } + defer func() { _ = client.Stop() }() + + stream, err := client.OpenStream(context.Background(), StreamOpenOptions{Channel: StreamDataChannel}) + if err != nil { + t.Fatalf("client OpenStream failed: %v", err) + } + readStreamExactly(t, stream, "early", 2*time.Second) + _ = stream.Close() +} + +func TestServerRejectsQueuedStreamOpenFromStaleTransport(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + var handlerCalls atomic.Int32 + server.SetStreamHandler(func(StreamAcceptInfo) error { + handlerCalls.Add(1) + return nil + }) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical := server.bootstrapAcceptedLogical("stale-inbound-stream-open", nil, firstLeft) + if logical == nil { + t.Fatal("bootstrapAcceptedLogical should return logical") + } + staleTransport := logical.CurrentTransportConn() + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { + t.Fatalf("attach replacement transport: %v", err) + } + payload, err := encode(StreamOpenRequest{StreamID: "queued-stale-open", DataID: 1}) + if err != nil { + t.Fatalf("encode StreamOpenRequest: %v", err) + } + message := Message{ + NetType: NET_SERVER, + LogicalConn: logical, + TransportConn: staleTransport, + TransferMsg: TransferMsg{ + Key: StreamOpenSignalKey, + Value: payload, + Type: MSG_ASYNC, + }, + } + + server.handleInboundStreamOpen(&message) + if got := handlerCalls.Load(); got != 0 { + t.Fatalf("stale StreamOpen handler calls = %d, want 0", got) + } + if stream, ok := server.getStreamRuntime().lookup(serverFileScope(logical), "queued-stale-open"); ok { + t.Fatalf("stale StreamOpen registered runtime handle: %+v", stream.snapshot()) + } +} + +func TestClientRejectsStreamOpenWhenRouteReattachesDuringRegistration(t *testing.T) { + client := NewClient().(*ClientCommon) + UseLegacySecurityClient(client) + var handlerCalls atomic.Int32 + client.SetStreamHandler(func(StreamAcceptInfo) error { + handlerCalls.Add(1) + return nil + }) + + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + epoch := client.beginClientSessionEpoch() + client.setClientSessionRuntime(newClientSessionRuntime(firstLeft, stopCtx, stopFn, queue, epoch)) + client.markSessionStarted() + defer client.markSessionStopped("test done", nil) + + payload, err := encode(StreamOpenRequest{StreamID: "client-stale-inbound-open", DataID: 2}) + if err != nil { + t.Fatalf("encode StreamOpenRequest: %v", err) + } + message := Message{ + NetType: NET_CLIENT, + ServerConn: client, + clientRoute: client.clientSessionRouteSnapshot(), + TransferMsg: TransferMsg{ + Key: StreamOpenSignalKey, + Value: payload, + Type: MSG_ASYNC, + }, + } + + runtime := client.getStreamRuntime() + runtime.mu.Lock() + done := make(chan struct{}) + go func() { + defer close(done) + client.handleInboundStreamOpen(&message) + }() + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := client.attachClientSessionTransport(secondLeft); err != nil { + runtime.mu.Unlock() + t.Fatalf("attach client replacement transport: %v", err) + } + runtime.mu.Unlock() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("client stale StreamOpen handler did not return") + } + if got := handlerCalls.Load(); got != 0 { + t.Fatalf("stale client StreamOpen handler calls = %d, want 0", got) + } + if stream, ok := runtime.lookup(clientFileScope(), "client-stale-inbound-open"); ok { + t.Fatalf("stale client StreamOpen registered runtime handle: %+v", stream.snapshot()) + } +} + +func TestClientStaleStreamRouteCannotAffectReplacement(t *testing.T) { + client := NewClient().(*ClientCommon) + runtime := client.getStreamRuntime() + epoch := client.beginClientSessionEpoch() + + staleLeft, staleRight := net.Pipe() + defer staleLeft.Close() + defer staleRight.Close() + currentLeft, currentRight := net.Pipe() + defer currentLeft.Close() + defer currentRight.Close() + staleRoute := clientSessionRoute{binding: newTransportBinding(staleLeft, nil), epoch: epoch} + currentRoute := clientSessionRoute{binding: newTransportBinding(currentLeft, nil), epoch: epoch} + + stream := newStreamHandle(context.Background(), runtime, clientFileScope(), StreamOpenRequest{ + StreamID: "replacement-stream", + DataID: 2, + Channel: StreamDataChannel, + }, epoch, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) + stream.setClientSessionRoute(currentRoute) + if err := runtime.register(clientFileScope(), stream); err != nil { + t.Fatalf("register replacement stream: %v", err) + } + + closePayload, err := encode(StreamCloseRequest{StreamID: stream.ID(), Full: true}) + if err != nil { + t.Fatalf("encode close request: %v", err) + } + client.handleInboundStreamClose(&Message{ + NetType: NET_CLIENT, + ServerConn: client, + clientRoute: staleRoute, + TransferMsg: TransferMsg{Key: StreamCloseSignalKey, Value: closePayload, Type: MSG_ASYNC}, + }) + + resetPayload, err := encode(StreamResetRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot(), Error: "stale reset"}) + if err != nil { + t.Fatalf("encode reset request: %v", err) + } + client.handleInboundStreamReset(&Message{ + NetType: NET_CLIENT, + ServerConn: client, + clientRoute: staleRoute, + TransferMsg: TransferMsg{Key: StreamResetSignalKey, Value: resetPayload, Type: MSG_ASYNC}, + }) + mismatchedClose, err := encode(StreamCloseRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot() + 2, Full: true}) + if err != nil { + t.Fatalf("encode mismatched close request: %v", err) + } + client.handleInboundStreamClose(&Message{ + NetType: NET_CLIENT, + ServerConn: client, + clientRoute: currentRoute, + TransferMsg: TransferMsg{Key: StreamCloseSignalKey, Value: mismatchedClose, Type: MSG_ASYNC}, + }) + mismatchedReset, err := encode(StreamResetRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot() + 2, Error: "mismatched reset"}) + if err != nil { + t.Fatalf("encode mismatched reset request: %v", err) + } + client.handleInboundStreamReset(&Message{ + NetType: NET_CLIENT, + ServerConn: client, + clientRoute: currentRoute, + TransferMsg: TransferMsg{Key: StreamResetSignalKey, Value: mismatchedReset, Type: MSG_ASYNC}, + }) + client.dispatchStreamEnvelopeAtRoute(staleRoute, newStreamDataEnvelope(stream.ID(), []byte("stale-envelope"))) + client.dispatchFastStreamDataWithOwnerAtRoute(staleRoute, streamFastDataFrame{ + DataID: stream.dataIDSnapshot(), + Seq: 1, + Payload: []byte("stale-fast"), + }, nil) + + stream.mu.Lock() + defer stream.mu.Unlock() + if stream.remoteClosed || stream.peerReadClosed || stream.resetErr != nil || len(stream.readQueue) != 0 || len(stream.readBuf.data) != 0 { + t.Fatalf("stale route mutated replacement: remoteClosed=%v peerReadClosed=%v reset=%v queued=%d buffered=%d", stream.remoteClosed, stream.peerReadClosed, stream.resetErr, len(stream.readQueue), len(stream.readBuf.data)) + } +} + +func TestServerStaleStreamControlsCannotAffectReplacement(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical := server.bootstrapAcceptedLogical("stale-stream-control", nil, firstLeft) + if logical == nil { + t.Fatal("bootstrapAcceptedLogical should return logical") + } + staleTransport := logical.CurrentTransportConn() + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { + t.Fatalf("attach replacement transport: %v", err) + } + currentTransport := logical.CurrentTransportConn() + if currentTransport == nil || currentTransport == staleTransport { + t.Fatal("replacement transport should be current") + } + + runtime := server.getStreamRuntime() + scope := serverFileScope(logical) + stream := newStreamHandle(logical.stopContextSnapshot(), runtime, scope, StreamOpenRequest{ + StreamID: "replacement-stream", + DataID: 1, + Channel: StreamDataChannel, + }, 0, logical, currentTransport, currentTransport.TransportGeneration(), nil, nil, nil, runtime.configSnapshot()) + if err := runtime.register(scope, stream); err != nil { + t.Fatalf("register replacement stream: %v", err) + } + + closePayload, err := encode(StreamCloseRequest{StreamID: stream.ID(), Full: true}) + if err != nil { + t.Fatalf("encode close request: %v", err) + } + server.handleInboundStreamClose(&Message{ + NetType: NET_SERVER, + LogicalConn: logical, + TransportConn: staleTransport, + TransferMsg: TransferMsg{Key: StreamCloseSignalKey, Value: closePayload, Type: MSG_ASYNC}, + }) + + resetPayload, err := encode(StreamResetRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot(), Error: "stale reset"}) + if err != nil { + t.Fatalf("encode reset request: %v", err) + } + server.handleInboundStreamReset(&Message{ + NetType: NET_SERVER, + LogicalConn: logical, + TransportConn: staleTransport, + TransferMsg: TransferMsg{Key: StreamResetSignalKey, Value: resetPayload, Type: MSG_ASYNC}, + }) + + stream.mu.Lock() + defer stream.mu.Unlock() + if stream.remoteClosed || stream.peerReadClosed || stream.resetErr != nil { + t.Fatalf("stale controls mutated replacement: remoteClosed=%v peerReadClosed=%v reset=%v", stream.remoteClosed, stream.peerReadClosed, stream.resetErr) + } +} diff --git a/stream_runtime.go b/stream_runtime.go index a530a29..e11de91 100644 --- a/stream_runtime.go +++ b/stream_runtime.go @@ -9,24 +9,38 @@ import ( ) type streamRuntime struct { - rolePrefix string - seq atomic.Uint64 - dataSeq atomic.Uint64 + rolePrefix string + seq atomic.Uint64 + dataSeq uint64 + peerDataSeq uint64 + dataStart uint64 + dataStep uint64 - mu sync.RWMutex - handler func(StreamAcceptInfo) error - streams map[string]*streamHandle - data map[string]map[uint64]*streamHandle - cfg streamConfig - flow *streamFlowController + mu sync.RWMutex + handler func(StreamAcceptInfo) error + streams map[string]*streamHandle + data map[string]map[uint64]*streamHandle + reserved map[string]map[uint64]struct{} + cfg streamConfig + flow *streamFlowController } func newStreamRuntime(rolePrefix string) *streamRuntime { cfg := defaultStreamConfig() + dataStart, dataStep := uint64(1), uint64(1) + if rolePrefix == "cstrm" { + dataStep = 2 + } else if rolePrefix == "sstrm" { + dataStart = 2 + dataStep = 2 + } return &streamRuntime{ rolePrefix: rolePrefix, + dataStart: dataStart, + dataStep: dataStep, streams: make(map[string]*streamHandle), data: make(map[string]map[uint64]*streamHandle), + reserved: make(map[string]map[uint64]struct{}), cfg: cfg, flow: newStreamFlowController(cfg), } @@ -43,7 +57,123 @@ func (r *streamRuntime) nextDataID() uint64 { if r == nil { return 0 } - return r.dataSeq.Add(1) + r.mu.Lock() + defer r.mu.Unlock() + id, _ := r.nextDataIDLocked(defaultFileScope, false) + return id +} + +func (r *streamRuntime) reserveDataID(scope string) (uint64, error) { + if r == nil { + return 0, errStreamRuntimeNil + } + scope = normalizeFileScope(scope) + r.mu.Lock() + defer r.mu.Unlock() + id, err := r.nextDataIDLocked(scope, true) + return id, err +} + +func (r *streamRuntime) releaseDataID(scope string, dataID uint64) { + if r == nil || dataID == 0 { + return + } + scope = normalizeFileScope(scope) + r.mu.Lock() + defer r.mu.Unlock() + if reserved := r.reserved[scope]; reserved != nil { + delete(reserved, dataID) + if len(reserved) == 0 { + delete(r.reserved, scope) + } + } +} + +func (r *streamRuntime) nextDataIDLocked(scope string, reserve bool) (uint64, error) { + if r == nil { + return 0, errStreamRuntimeNil + } + for { + candidate, ok := r.nextDataCandidateLocked(false) + if !ok { + return 0, errStreamDataIDExhausted + } + if r.dataInUseLocked(scope, candidate) { + continue + } + if reserve { + reserved := r.reserved[scope] + if reserved == nil { + reserved = make(map[uint64]struct{}) + r.reserved[scope] = reserved + } + reserved[candidate] = struct{}{} + } + return candidate, nil + } +} + +func (r *streamRuntime) nextPeerDataIDLocked(scope string) (uint64, error) { + if r == nil { + return 0, errStreamRuntimeNil + } + for { + candidate, ok := r.nextDataCandidateLocked(true) + if !ok { + return 0, errStreamDataIDExhausted + } + if r.dataInUseLocked(scope, candidate) { + continue + } + return candidate, nil + } +} + +func (r *streamRuntime) nextDataCandidateLocked(peer bool) (uint64, bool) { + seq := &r.dataSeq + start, step := r.dataStart, r.dataStep + if peer && step == 2 { + seq = &r.peerDataSeq + start = 3 - r.dataStart + } + if step == 0 { + step = 1 + } + if *seq == 0 { + if start == 0 { + return 0, false + } + *seq = start + return start, true + } + if *seq > ^uint64(0)-step { + return 0, false + } + candidate := *seq + step + if step == 2 && candidate%2 != start%2 { + if candidate == ^uint64(0) { + return 0, false + } + candidate++ + } + *seq = candidate + return candidate, true +} + +func (r *streamRuntime) dataInUseLocked(scope string, dataID uint64) bool { + if dataID == 0 { + return true + } + if dataScope := r.data[scope]; dataScope != nil { + if _, ok := dataScope[dataID]; ok { + return true + } + } + if reserved := r.reserved[scope]; reserved != nil { + _, ok := reserved[dataID] + return ok + } + return false } func (r *streamRuntime) setHandler(fn func(StreamAcceptInfo) error) { @@ -65,6 +195,18 @@ func (r *streamRuntime) handlerSnapshot() func(StreamAcceptInfo) error { } func (r *streamRuntime) register(scope string, stream *streamHandle) error { + return r.registerWithDirection(scope, stream, false, false) +} + +func (r *streamRuntime) registerInbound(scope string, stream *streamHandle) error { + return r.registerWithDirection(scope, stream, true, false) +} + +func (r *streamRuntime) registerReserved(scope string, stream *streamHandle) error { + return r.registerWithDirection(scope, stream, false, true) +} + +func (r *streamRuntime) registerWithDirection(scope string, stream *streamHandle, inbound bool, consumeReservation bool) error { if r == nil { return errStreamRuntimeNil } @@ -78,6 +220,17 @@ func (r *streamRuntime) register(scope string, stream *streamHandle) error { if _, ok := r.streams[key]; ok { return errStreamAlreadyExists } + if stream.dataID == 0 { + var err error + if inbound { + stream.dataID, err = r.nextPeerDataIDLocked(scope) + } else { + stream.dataID, err = r.nextDataIDLocked(scope, false) + } + if err != nil { + return err + } + } if stream.dataID != 0 { dataScope := r.data[scope] if dataScope == nil { @@ -87,12 +240,50 @@ func (r *streamRuntime) register(scope string, stream *streamHandle) error { if _, ok := dataScope[stream.dataID]; ok { return errStreamAlreadyExists } + if reserved := r.reserved[scope]; reserved != nil { + if _, ok := reserved[stream.dataID]; ok { + if !consumeReservation { + return errStreamAlreadyExists + } + delete(reserved, stream.dataID) + if len(reserved) == 0 { + delete(r.reserved, scope) + } + } + } dataScope[stream.dataID] = stream } r.streams[key] = stream return nil } +// adopt transfers ownership of a newly-created stream to the runtime. A +// failed registration is terminal so its child context cannot remain attached +// to the session after the caller drops the handle. +func (r *streamRuntime) adopt(scope string, stream *streamHandle) error { + err := r.register(scope, stream) + if err != nil && stream != nil { + stream.markReset(err) + } + return err +} + +func (r *streamRuntime) adoptInbound(scope string, stream *streamHandle) error { + err := r.registerInbound(scope, stream) + if err != nil && stream != nil { + stream.markReset(err) + } + return err +} + +func (r *streamRuntime) adoptReserved(scope string, stream *streamHandle) error { + err := r.registerReserved(scope, stream) + if err != nil && stream != nil { + stream.markReset(err) + } + return err +} + func (r *streamRuntime) lookup(scope string, streamID string) (*streamHandle, bool) { if r == nil || streamID == "" { return nil, false @@ -119,17 +310,47 @@ func (r *streamRuntime) lookupByDataID(scope string, dataID uint64) (*streamHand return stream, ok } -func (r *streamRuntime) remove(scope string, streamID string) { - if r == nil || streamID == "" { +func (r *streamRuntime) lookupControl(scope string, streamID string, dataID uint64) (*streamHandle, bool) { + if r == nil { + return nil, false + } + scope = normalizeFileScope(scope) + r.mu.RLock() + defer r.mu.RUnlock() + if streamID != "" { + stream, ok := r.streams[streamRuntimeKey(scope, streamID)] + if !ok || stream == nil { + return nil, false + } + if dataID != 0 && stream.dataID != dataID { + return nil, false + } + return stream, true + } + if dataID == 0 { + return nil, false + } + stream := r.data[scope][dataID] + return stream, stream != nil +} + +func (r *streamRuntime) remove(scope string, expected *streamHandle) { + if r == nil || expected == nil || expected.id == "" { return } scope = normalizeFileScope(scope) - key := streamRuntimeKey(scope, streamID) + key := streamRuntimeKey(scope, expected.id) r.mu.Lock() defer r.mu.Unlock() - if stream := r.streams[key]; stream != nil && stream.dataID != 0 { + stream := r.streams[key] + if stream != expected { + return + } + if stream.dataID != 0 { if dataScope := r.data[scope]; dataScope != nil { - delete(dataScope, stream.dataID) + if dataScope[stream.dataID] == stream { + delete(dataScope, stream.dataID) + } if len(dataScope) == 0 { delete(r.data, scope) } @@ -187,6 +408,44 @@ func (r *streamRuntime) closeScope(scope string, err error) { }, err) } +func (r *streamRuntime) closeClientRoute(route clientSessionRoute, err error) { + if r == nil { + return + } + if !r.mu.TryRLock() { + go r.closeClientRouteBlocking(route, err) + return + } + streams := r.collectClientRouteLocked(route) + r.mu.RUnlock() + r.resetClientRouteHandles(streams, err) +} + +func (r *streamRuntime) closeClientRouteBlocking(route clientSessionRoute, err error) { + r.mu.RLock() + streams := r.collectClientRouteLocked(route) + r.mu.RUnlock() + r.resetClientRouteHandles(streams, err) +} + +func (r *streamRuntime) collectClientRouteLocked(route clientSessionRoute) []*streamHandle { + streams := make([]*streamHandle, 0) + for _, stream := range r.streams { + if stream == nil || !sameClientSessionRoute(stream.clientRoute, route) { + continue + } + streams = append(streams, stream) + } + return streams +} + +func (r *streamRuntime) resetClientRouteHandles(streams []*streamHandle, err error) { + resetErr := streamRuntimeCloseError(err) + for _, stream := range streams { + stream.markReset(resetErr) + } +} + func (r *streamRuntime) closeMatching(match func(string) bool, err error) { if r == nil || match == nil { return diff --git a/stream_shared_batch.go b/stream_shared_batch.go index 29aab79..f97d0ef 100644 --- a/stream_shared_batch.go +++ b/stream_shared_batch.go @@ -47,11 +47,30 @@ func streamFastBatchPlainLen(frames []streamFastDataFrame) int { return total } -func encodeStreamFastBatchPlain(frames []streamFastDataFrame) ([]byte, error) { - if len(frames) == 0 { - return nil, errStreamFastPayloadInvalid +func streamFastBatchPlainLenChecked(frames []streamFastDataFrame) (int, error) { + if len(frames) == 0 || len(frames) > streamFastBatchMaxItems { + return 0, errStreamFastPayloadInvalid } - buf := make([]byte, streamFastBatchPlainLen(frames)) + total := streamFastBatchHeaderLen + for _, frame := range frames { + itemLen := streamFastBatchFrameLen(frame) + if itemLen < streamFastBatchItemHeaderLen || itemLen > streamFastBatchMaxPlainBytes { + return 0, errStreamFastPayloadInvalid + } + if total > streamFastBatchMaxPlainBytes-itemLen { + return 0, errStreamFastPayloadInvalid + } + total += itemLen + } + return total, nil +} + +func encodeStreamFastBatchPlain(frames []streamFastDataFrame) ([]byte, error) { + plainLen, err := streamFastBatchPlainLenChecked(frames) + if err != nil { + return nil, err + } + buf := make([]byte, plainLen) if err := writeStreamFastBatchPlain(buf, frames); err != nil { return nil, err } @@ -62,14 +81,21 @@ func encodeStreamFastBatchPayloadFast(encode transportFastPlainEncoder, secretKe if encode == nil { return nil, errTransportPayloadEncryptFailed } - plainLen := streamFastBatchPlainLen(frames) + plainLen, err := streamFastBatchPlainLenChecked(frames) + if err != nil { + return nil, err + } return encode(secretKey, plainLen, func(dst []byte) error { return writeStreamFastBatchPlain(dst, frames) }) } func writeStreamFastBatchPlain(dst []byte, frames []streamFastDataFrame) error { - if len(frames) == 0 || len(dst) != streamFastBatchPlainLen(frames) { + plainLen, err := streamFastBatchPlainLenChecked(frames) + if err != nil { + return err + } + if len(dst) != plainLen { return errStreamFastPayloadInvalid } copy(dst[:4], streamFastBatchMagic) @@ -101,10 +127,11 @@ func walkStreamFastBatchPlain(payload []byte, fn func(streamFastDataFrame) error if payload[4] != streamFastBatchVersion { return true, errStreamFastPayloadInvalid } - count := int(binary.BigEndian.Uint32(payload[8:12])) - if count <= 0 { + wireCount := binary.BigEndian.Uint32(payload[8:12]) + if wireCount == 0 || wireCount > streamFastBatchMaxItems { return true, errStreamFastPayloadInvalid } + count := int(wireCount) offset := streamFastBatchHeaderLen for index := 0; index < count; index++ { if len(payload)-offset < streamFastBatchItemHeaderLen { @@ -113,11 +140,12 @@ func walkStreamFastBatchPlain(payload []byte, fn func(streamFastDataFrame) error flags := payload[offset] dataID := binary.BigEndian.Uint64(payload[offset+4 : offset+12]) seq := binary.BigEndian.Uint64(payload[offset+12 : offset+20]) - payloadLen := int(binary.BigEndian.Uint32(payload[offset+20 : offset+24])) + wirePayloadLen := binary.BigEndian.Uint32(payload[offset+20 : offset+24]) offset += streamFastBatchItemHeaderLen - if dataID == 0 || payloadLen < 0 || len(payload)-offset < payloadLen { + if dataID == 0 || uint64(wirePayloadLen) > uint64(len(payload)-offset) { return true, errStreamFastPayloadInvalid } + payloadLen := int(wirePayloadLen) if fn != nil { if err := fn(streamFastDataFrame{ Flags: flags, diff --git a/transfer_observability_test.go b/transfer_observability_test.go index 0a3c105..716d2b8 100644 --- a/transfer_observability_test.go +++ b/transfer_observability_test.go @@ -39,6 +39,7 @@ func (s *transferDelayedWriteStream) Write(p []byte) (int, error) { type transferDelayedCommitSink struct { data []byte writeDelay time.Duration + readDelay time.Duration syncDelay time.Duration commitDelay time.Duration } @@ -65,6 +66,7 @@ func (s *transferDelayedCommitSink) WriteAt(p []byte, off int64) (int, error) { } func (s *transferDelayedCommitSink) ReadAt(p []byte, off int64) (int, error) { + time.Sleep(s.readDelay) if off < 0 || off >= int64(len(s.data)) { return 0, io.EOF } @@ -155,6 +157,7 @@ func TestTransferReceiveSessionCommitRecordsTelemetry(t *testing.T) { const ( writeDelay = 4 * time.Millisecond syncDelay = 3 * time.Millisecond + verifyDelay = 5 * time.Millisecond commitDelay = 5 * time.Millisecond ) data := []byte("abcdefgh") @@ -169,6 +172,7 @@ func TestTransferReceiveSessionCommitRecordsTelemetry(t *testing.T) { }) sink := newTransferDelayedCommitSink(len(data), writeDelay, syncDelay, commitDelay) + sink.readDelay = verifyDelay session := newTransferReceiveSession(scope, scope, nil, nil, 0, TransferReceiveOptions{ Descriptor: TransferDescriptor{ ID: transferID, @@ -201,8 +205,8 @@ func TestTransferReceiveSessionCommitRecordsTelemetry(t *testing.T) { if got := snapshot.SyncDuration; got < 3*syncDelay { t.Fatalf("sync duration = %v, want at least %v", got, 3*syncDelay) } - if got := snapshot.VerifyDuration; got <= 0 { - t.Fatalf("verify duration = %v, want > 0", got) + if got := snapshot.VerifyDuration; got < verifyDelay { + t.Fatalf("verify duration = %v, want at least %v", got, verifyDelay) } if got := snapshot.CommitDuration; got < commitDelay { t.Fatalf("commit duration = %v, want at least %v", got, commitDelay) diff --git a/transfer_plane.go b/transfer_plane.go index c450a1f..310a3df 100644 --- a/transfer_plane.go +++ b/transfer_plane.go @@ -18,12 +18,15 @@ const ( transferStreamMetadataKindKey = "_notify.transfer_stream_kind" transferStreamMetadataKindValue = "segment" transferFrameHeaderSize = 4 + transferFrameMaxPayloadBytes = transportFrameMaxPayloadBytes transferFrameAggregateLimit = 128 * 1024 transferFrameAggregateCount = 8 transferCommitWaitTimeout = 30 * time.Second transferChecksumChunkSize = 64 * 1024 ) +var errTransferFrameTooLarge = errors.New("transfer frame too large") + type transferSendTarget struct { runtime *transferRuntime runtimeScope string @@ -843,6 +846,9 @@ func readTransferFrame(stream Stream) ([]byte, error) { return nil, err } length := binary.BigEndian.Uint32(header) + if length > transferFrameMaxPayloadBytes { + return nil, fmt.Errorf("%w: payload=%d max=%d", errTransferFrameTooLarge, length, transferFrameMaxPayloadBytes) + } payload := make([]byte, int(length)) if _, err := io.ReadFull(stream, payload); err != nil { return nil, err diff --git a/transfer_send_pipeline.go b/transfer_send_pipeline.go index 9e3394d..f4c6e52 100644 --- a/transfer_send_pipeline.go +++ b/transfer_send_pipeline.go @@ -3,6 +3,7 @@ package notify import ( "context" "errors" + "fmt" "io" "time" @@ -32,6 +33,9 @@ func (w *transferFrameBatchWriter) writeEncodedFrame(payload []byte) error { if w == nil { return nil } + if len(payload) > transferFrameMaxPayloadBytes { + return fmt.Errorf("%w: payload=%d max=%d", errTransferFrameTooLarge, len(payload), transferFrameMaxPayloadBytes) + } frame := buildTransferFrame(payload) if len(w.batch) > 0 && len(w.batch)+len(frame) > transferFrameAggregateLimit { if err := w.flush(); err != nil { diff --git a/transport_codec.go b/transport_codec.go index b0ad87f..38b44ba 100644 --- a/transport_codec.go +++ b/transport_codec.go @@ -112,6 +112,9 @@ func (c *ClientCommon) encodeTransferMsg(msg TransferMsg) ([]byte, error) { if queue == nil { return nil, errClientSessionQueueUnavailable } + if err := validateTransportFramePayloadLen(data); err != nil { + return nil, err + } return queue.BuildMessage(data), nil } @@ -151,6 +154,9 @@ func (s *ServerCommon) encodeTransferMsg(c *ClientConn, msg TransferMsg) ([]byte if queue == nil { return nil, errServerSessionQueueUnavailable } + if err := validateTransportFramePayloadLen(data); err != nil { + return nil, err + } return queue.BuildMessage(data), nil } @@ -206,6 +212,9 @@ func (c *ClientCommon) encodeEnvelope(env Envelope) ([]byte, error) { if queue == nil { return nil, errClientSessionQueueUnavailable } + if err := validateTransportFramePayloadLen(data); err != nil { + return nil, err + } return queue.BuildMessage(data), nil } @@ -286,6 +295,9 @@ func (s *ServerCommon) encodeEnvelopeLogical(logical *LogicalConn, env Envelope) if queue == nil { return nil, errServerSessionQueueUnavailable } + if err := validateTransportFramePayloadLen(data); err != nil { + return nil, err + } return queue.BuildMessage(data), nil } diff --git a/transport_conn.go b/transport_conn.go index 439083e..8933281 100644 --- a/transport_conn.go +++ b/transport_conn.go @@ -9,9 +9,14 @@ import ( ) type TransportConn struct { - logical *LogicalConn - generation uint64 - remoteAddr net.Addr + logical *LogicalConn + generation uint64 + remoteAddr net.Addr + // binding pins this view to the physical connection that produced it. A + // logical session may replace its transport while an operation is in + // flight; using the logical's current binding at that point would cross the + // reconnect boundary. + binding *transportBinding attached bool hasRuntimeConn bool } @@ -19,6 +24,7 @@ type TransportConn struct { const ( transportStreamReadBufferSize = 1024 * 1024 transportPacketReadBufferSize = 64 * 1024 + transportFrameMaxPayloadBytes = 64 * 1024 * 1024 ) func streamReadBuffer() []byte { @@ -120,6 +126,7 @@ func (c *LogicalConn) currentTransportConnSnapshot() *TransportConn { logical: logical, generation: c.transportGenerationSnapshot(), remoteAddr: remoteAddr, + binding: c.transportBindingSnapshot(), attached: true, hasRuntimeConn: hasRuntimeConn, } @@ -132,6 +139,7 @@ func (c *LogicalConn) currentTransportConnSnapshot() *TransportConn { logical: logical, generation: c.transportGenerationSnapshot(), remoteAddr: remoteAddr, + binding: c.transportBindingSnapshot(), attached: true, hasRuntimeConn: hasRuntimeConn, } @@ -197,13 +205,42 @@ func (t *TransportConn) IsCurrent() bool { current := logical.CurrentTransportConn() if current == nil { return false - } - if current.generation != t.generation { + } else if current.generation != t.generation { return false + } else if t.binding != nil { + return current.binding == t.binding } return transportConnAddrString(current.remoteAddr) == transportConnAddrString(t.remoteAddr) } +// serverTransportBindingSnapshot returns the physical binding pinned by an +// explicit transport. The logical fallback preserves compatibility for +// package-local callers that construct TransportConn values directly. +func serverTransportBindingSnapshot(logical *LogicalConn, transport *TransportConn) *transportBinding { + if transport != nil && transport.binding != nil { + return transport.binding + } + if logical == nil { + return nil + } + return logical.transportBindingSnapshot() +} + +// serverTransportBindingSnapshotForConn preserves an inbound socket handoff: +// peer attach can transfer the same physical conn from a temporary logical +// peer to its stable logical peer. In that case the destination's binding owns +// the exact inbound conn and must serialize the reply. A different replacement +// conn can never satisfy this identity check, so explicit transport sends stay +// pinned to their original binding. +func serverTransportBindingSnapshotForConn(logical *LogicalConn, transport *TransportConn, conn net.Conn) *transportBinding { + if conn != nil && logical != nil { + if current := logical.transportBindingSnapshot(); current != nil && current.connSnapshot() == conn { + return current + } + } + return serverTransportBindingSnapshot(logical, transport) +} + func transportConnAddrString(addr net.Addr) string { if addr == nil { return "" diff --git a/transport_conn_test.go b/transport_conn_test.go index 9ecdbd5..271435f 100644 --- a/transport_conn_test.go +++ b/transport_conn_test.go @@ -171,6 +171,49 @@ func TestTransportConnSendRejectsStaleGenerationAfterReattach(t *testing.T) { } } +func TestTransportConnPinnedBindingDoesNotWriteReplacementTransport(t *testing.T) { + server := NewServer().(*ServerCommon) + UseLegacySecurityServer(server) + + runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) + defer runtimeCancel() + queue := stario.NewQueueCtx(runtimeCtx, 4, math.MaxUint32) + server.setServerSessionRuntime(&serverSessionRuntime{stopCtx: runtimeCtx, stopFn: runtimeCancel, queue: queue}) + server.markSessionStarted() + defer server.markSessionStopped("test done", nil) + + firstLeft, firstRight := net.Pipe() + defer firstRight.Close() + logical, _, _ := newRegisteredServerLogicalForTest(t, server, "transport-pinned-binding", firstLeft, runtimeCtx, runtimeCancel) + firstTransport := logical.CurrentTransportConn() + if firstTransport == nil || firstTransport.binding == nil { + t.Fatal("first transport should pin its physical binding") + } + + secondLeft, secondRight := net.Pipe() + defer secondRight.Close() + if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { + t.Fatalf("attach replacement transport: %v", err) + } + if got := serverTransportBindingSnapshot(logical, firstTransport); got != firstTransport.binding { + t.Fatal("stale explicit transport resolved to a different physical binding") + } + + result := make(chan error, 1) + go func() { + result <- server.writeEnvelopePayloadContext(context.Background(), logical, firstTransport, nil, []byte("stale")) + }() + assertNoPipeWrite(t, secondRight, "pinned transport write crossed onto replacement transport") + select { + case err := <-result: + if err == nil { + t.Fatal("write to retired pinned transport unexpectedly succeeded") + } + case <-time.After(time.Second): + t.Fatal("write to retired pinned transport did not terminate") + } +} + func TestTransportConnRuntimeSnapshotIncludesDetachDiagnostics(t *testing.T) { server := NewServer().(*ServerCommon) left, right := net.Pipe() diff --git a/transport_write.go b/transport_write.go index 7dd3f04..b610460 100644 --- a/transport_write.go +++ b/transport_write.go @@ -4,6 +4,7 @@ import ( "b612.me/stario" "context" "errors" + "fmt" "io" "net" "strings" @@ -14,6 +15,13 @@ import ( var transportConnWriteGates sync.Map var errTransportFrameQueueUnavailable = errors.New("transport frame queue is unavailable") +func validateTransportFramePayloadLen(payload []byte) error { + if len(payload) > transportFrameMaxPayloadBytes { + return fmt.Errorf("%w: %d > %d", stario.ErrQueueMessageTooLarge, len(payload), transportFrameMaxPayloadBytes) + } + return nil +} + type connWriteGateRef struct { mu sync.Mutex gate chan struct{} @@ -270,6 +278,9 @@ func writeFramedPayloadUnlocked(conn net.Conn, queue *stario.StarQueue, payload if queue == nil { return errTransportFrameQueueUnavailable } + if err := validateTransportFramePayloadLen(payload); err != nil { + return err + } if isPacketTransportConn(conn) { return writeFullToConnUnlocked(conn, queue.BuildMessage(payload)) } @@ -286,6 +297,11 @@ func writeFramedPayloadBatchUnlocked(conn net.Conn, queue *stario.StarQueue, pay if len(payloads) == 0 { return nil } + for _, payload := range payloads { + if err := validateTransportFramePayloadLen(payload); err != nil { + return err + } + } if isPacketTransportConn(conn) { for _, payload := range payloads { if err := writeFullToConnUnlocked(conn, queue.BuildMessage(payload)); err != nil { diff --git a/transport_write_test.go b/transport_write_test.go index ecc8cf5..70f7850 100644 --- a/transport_write_test.go +++ b/transport_write_test.go @@ -439,6 +439,127 @@ func TestStreamBatchSenderCarriesContextDeadlineIntoPhysicalWrite(t *testing.T) } } +func TestBatchSenderQueueWaitTimeoutDoesNotPoisonSender(t *testing.T) { + tests := []struct { + name string + new func(*transportBinding) batchSenderTestSender + }{ + { + name: "bulk", + new: func(binding *transportBinding) batchSenderTestSender { + sender := newTestBulkBatchSender(binding) + return batchSenderTestAdapter{ + submitFn: func(ctx context.Context) error { + return sender.submitData(ctx, 1, 1, bulkFastPathVersionV1, []byte("bulk")) + }, + errFn: sender.errSnapshot, + stopFn: sender.stop, + } + }, + }, + { + name: "stream", + new: func(binding *transportBinding) batchSenderTestSender { + sender := newTestStreamBatchSender(binding, nil) + return batchSenderTestAdapter{ + submitFn: func(ctx context.Context) error { + return sender.submitData(ctx, 1, 1, streamFastPathVersionV1, []byte("stream")) + }, + errFn: sender.errSnapshot, + stopFn: sender.stop, + } + }, + }, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + binding := newTransportBinding(&serializedWriteTestConn{}, stario.NewQueue()) + sender := tc.new(binding) + if err := binding.lockConnWriteContext(context.Background()); err != nil { + t.Fatalf("lock shared write gate: %v", err) + } + defer binding.unlockConnWrite() + ctx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond) + defer cancel() + err := sender.submit(ctx) + if !isBatchSenderQueueWaitError(err) || !isTimeoutLikeError(err) { + t.Fatalf("queue wait error=%v, want classified timeout", err) + } + if got := sender.errSnapshot(); got != nil { + t.Fatalf("caller queue timeout poisoned healthy sender: %v", got) + } + sender.stop() + }) + } +} + +type batchSenderTestSender interface { + submit(context.Context) error + errSnapshot() error + stop() +} + +type batchSenderTestAdapter struct { + submitFn func(context.Context) error + errFn func() error + stopFn func() +} + +func (a batchSenderTestAdapter) submit(ctx context.Context) error { return a.submitFn(ctx) } +func (a batchSenderTestAdapter) errSnapshot() error { return a.errFn() } +func (a batchSenderTestAdapter) stop() { a.stopFn() } + +func TestBatchSenderFailBatchCompletesEveryRequest(t *testing.T) { + tests := []struct { + name string + bulk bool + }{ + {name: "bulk", bulk: true}, + {name: "stream", bulk: false}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + var released atomic.Int32 + if tc.bulk { + sender := &bulkBatchSender{} + sender.queued.Store(2) + requests := []bulkBatchRequest{ + {done: make(chan error, 1), release: func() { released.Add(1) }}, + {done: make(chan error, 1), release: func() { released.Add(1) }}, + } + sender.failBatch(requests, errTransportDetached) + for i, req := range requests { + if err := <-req.done; !errors.Is(err, errTransportDetached) { + t.Fatalf("request %d error=%v, want transport detached", i, err) + } + } + if got := sender.queued.Load(); got != 0 { + t.Fatalf("queued=%d after local batch failure, want 0", got) + } + } else { + sender := &streamBatchSender{} + sender.queued.Store(2) + requests := []streamBatchRequest{ + {done: make(chan error, 1)}, + {done: make(chan error, 1)}, + } + sender.failBatch(requests, errTransportDetached) + for i, req := range requests { + if err := <-req.done; !errors.Is(err, errTransportDetached) { + t.Fatalf("request %d error=%v, want transport detached", i, err) + } + } + if got := sender.queued.Load(); got != 0 { + t.Fatalf("queued=%d after local batch failure, want 0", got) + } + } + if tc.bulk && released.Load() != 2 { + t.Fatalf("bulk payload releases=%d, want 2", released.Load()) + } + }) + } +} + func TestTransportBindingStopWithCloseInterruptsPhysicalWrite(t *testing.T) { left, right := net.Pipe() defer right.Close() @@ -889,6 +1010,30 @@ func TestControlBatchSenderCancelsQueuedRequestWithoutStoppingSender(t *testing. } } +func TestControlBatchWaitContextIsIndependentOfRequestCancellation(t *testing.T) { + sender := &controlBatchSender{stopCtx: context.Background()} + firstCtx, cancelFirst := context.WithCancel(context.Background()) + secondCtx, cancelSecond := context.WithCancel(context.Background()) + waitCtx, cleanup := sender.controlBatchWaitContext([]controlBatchRequest{ + {ctx: firstCtx}, + {ctx: secondCtx}, + }) + defer cleanup() + cancelFirst() + select { + case <-waitCtx.Done(): + t.Fatal("one canceled request canceled the physical batch wait") + case <-time.After(20 * time.Millisecond): + } + cancelSecond() + select { + case <-waitCtx.Done(): + case <-time.After(time.Second): + t.Fatal("physical batch wait did not cancel after all requests canceled") + } + cleanup() +} + func TestControlBatchSenderCancelsWhileWaitingForSharedWriteLock(t *testing.T) { binding := newTransportBinding(&serializedWriteTestConn{}, stario.NewQueue()) sender := newControlBatchSender(binding)