fix(notify): 修复传输生命周期竞态,完善背压与协议边界

- 完善 stream/bulk DataID 分配、预留和双向命名空间,修复并发打开及 dedicated/shared 回退时的 ID 冲突
- 将收发、回复、恢复任务和 sidecar 绑定原始会话与物理连接,防止重连后的旧消息误操作新连接
- 加强 close/reset 身份校验及实例移除检查,修复 dedicated attach 失败、通道引用和资源回收竞态
- 收紧批量发送器停止准入,确保在途入队完成后统一清理请求、缓冲区和等待者
- 修复 record 满队列死锁、取消时序号消耗及关闭竞态,确保关闭有界并返回真实错误
- 增加协商式 record 逻辑半关闭,保留反向 ACK;通过 reset 传递 RecordFailure,避免背压掩盖原始失败原因
- 补齐帧长度、批次数量、序号溢出和未确认窗口校验,提前拒绝超限数据并按字节预算拆批
- 为入站分发增加全局及单连接的条数、字节预算和阻塞背压,关闭时唤醒等待者,消除正常断连日志噪音
- 完善 bulk 窗口释放失败处理与传输诊断,补充并发、重连、背压、协议边界及真实 TCP 回归覆盖
This commit is contained in:
2026-09-23 15:33:17 +08:00
parent 0826e17063
commit 1f2e74acca
79 changed files with 9190 additions and 1013 deletions
+47
View File
@@ -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)
}
+299 -12
View File
@@ -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
@@ -318,6 +325,9 @@ type bulkHandle 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
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,6 +2144,7 @@ func (b *bulkHandle) finalize() {
if b == nil {
return
}
b.finalizeOnce.Do(func() {
b.markDedicatedAttachClosed()
b.maybeSendWindowRelease(0, true)
if b.cancel != nil {
@@ -1878,12 +2164,13 @@ func (b *bulkHandle) finalize() {
if b.client != nil && b.releaseDedicatedActiveReserved() {
b.client.releaseBulkDedicatedActiveSlot()
}
if b.client != nil {
b.client.releaseBulkDedicatedLane(b.dedicatedLaneIDSnapshot())
if b.client != nil && b.releaseDedicatedLaneReserved() {
b.client.releaseBulkDedicatedLaneAtRoute(b.dedicatedLaneIDSnapshot(), b.clientSessionRouteSnapshot())
}
if b.runtime != nil {
b.runtime.remove(b.runtimeScope, b.id)
b.runtime.remove(b.runtimeScope, b)
}
})
}
func (b *bulkHandle) recordReadLocked(n int, now time.Time) {
+39
View File
@@ -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")
}
})
}
}
+105 -22
View File
@@ -62,6 +62,9 @@ type bulkBatchSender struct {
doneCh chan struct{}
stopOnce sync.Once
admissionMu sync.Mutex
admitting sync.WaitGroup
admissionClosed bool
flushMu sync.Mutex
queued atomic.Int64
errMu sync.Mutex
@@ -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
+68
View File
@@ -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 {
+166 -42
View File
@@ -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,12 +213,9 @@ 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) {
_, 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
+457
View File
@@ -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()
}
+399 -46
View File
@@ -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)))
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 = routeWaitError(err)
return flightErr
}
defer releaseAttachSlot()
if sidecar := c.clientDedicatedSidecarSnapshotForLaneAtRoute(laneID, route); sidecar != nil {
if err := checkRoute(); err != nil {
flightErr = err
return err
}
defer releaseAttachSlot()
if sidecar := c.clientDedicatedSidecarSnapshotForLane(laneID); sidecar != nil && sidecar.conn != nil {
if err := bulk.attachDedicatedConnShared(sidecar.conn); err == nil {
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 {
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)
})
}
+245
View File
@@ -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)
+186 -41
View File
@@ -79,6 +79,9 @@ type bulkDedicatedSender struct {
stopCh chan struct{}
doneCh chan struct{}
stopOnce sync.Once
admissionMu sync.Mutex
admitting sync.WaitGroup
admissionClosed bool
flushMu sync.Mutex
queued atomic.Int64
@@ -220,20 +223,18 @@ func (s *bulkDedicatedSender) submitBatch(ctx context.Context, items []bulkDedic
req.Ack = make(chan error, 1)
}
s.queued.Add(1)
select {
case <-ctx.Done():
s.queued.Add(-1)
return normalizeStreamDeadlineError(ctx.Err())
case <-s.stopCh:
if !s.enqueue(req) {
s.queued.Add(-1)
if err := ctx.Err(); err != nil {
return normalizeStreamDeadlineError(err)
}
return s.stoppedErr()
case s.reqCh <- req:
}
if !wait {
return nil
}
return s.waitAck(req)
}
}
func (s *bulkDedicatedSender) tryDirectSubmitBatch(ctx context.Context, items []bulkDedicatedSendRequest) (bool, error) {
if s == nil {
@@ -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 {
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,
+121 -19
View File
@@ -2,6 +2,7 @@ package notify
import (
"context"
"fmt"
"net"
"sync"
"sync/atomic"
@@ -32,18 +33,23 @@ type bulkDedicatedLaneSender struct {
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,22 +241,19 @@ 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():
s.queued.Add(-1)
req.recycle()
return normalizeStreamDeadlineError(ctx.Err())
case <-s.stopCh:
if !s.enqueue(req) {
s.queued.Add(-1)
req.recycle()
if err := ctx.Err(); err != nil {
return normalizeStreamDeadlineError(err)
}
return s.stoppedErr()
case s.reqCh <- req:
}
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) {
if s == nil {
@@ -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 {
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,8 +748,20 @@ func (s *bulkDedicatedLaneSender) flush(batches []bulkDedicatedOutboundBatch, de
if release != nil {
defer release()
}
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) {
if s != nil {
@@ -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
+246 -37
View File
@@ -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)
})
}
+24 -9
View File
@@ -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,
+53 -23
View File
@@ -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
}
+183
View File
@@ -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
}
+279
View File
@@ -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)
}
}
+719
View File
@@ -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")
}
}
+428 -21
View File
@@ -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
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
}
if _, ok := dataScope[bulk.dataID]; ok {
bulk.dataID = dataID
} else if direction&bulkDataIndexOutbound != 0 && r.localDataID(bulk.dataID) && bulk.dataID > r.dataSeq {
r.dataSeq = bulk.dataID
}
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 {
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
+42 -10
View File
@@ -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,
+32
View File
@@ -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 {
+139
View File
@@ -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 {
+3
View File
@@ -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)
+169 -69
View File
@@ -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 {
if existing, exists := runtime.lookup(clientFileScope(), req.BulkID); exists {
if existing.acceptsClientSessionRoute(route) {
return nil, errBulkAlreadyExists
}
if req.Dedicated {
req.DedicatedLaneID = c.reserveBulkDedicatedLane()
existing.markReset(errTransportDetached)
}
if req.DataID == 0 {
req.DataID = runtime.nextDataID()
var reserveErr error
req.DataID, reserveErr = runtime.reserveDataID(clientFileScope(), 0)
if reserveErr != nil {
return nil, reserveErr
}
}
if req.Dedicated {
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{
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
}
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
}
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
}
return bulk, nil
}
resp, err := sendBulkOpenClient(ctx, c, req)
if err != nil {
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.DataID != 0 {
req.DataID = resp.DataID
}
if resp.FastPathVersion != 0 {
req.FastPathVersion = resp.FastPathVersion
bulk.setFastPathVersion(resp.FastPathVersion)
}
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
}
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(),
})
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
}
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,
+2
View File
@@ -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
}
+2 -1
View File
@@ -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):
+1 -1
View File
@@ -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):
+26 -8
View File
@@ -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)
}
+38 -10
View File
@@ -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,
+108
View File
@@ -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
}
+10
View File
@@ -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
+71
View File
@@ -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)
+65 -25
View File
@@ -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 {
if existing, exists := runtime.lookup(scope, req.StreamID); exists {
if existing.acceptsClientSessionRoute(route) {
return nil, errStreamAlreadyExists
}
resp, err := sendStreamOpenClient(ctx, c, req)
existing.markReset(transportDetachedSessionEpochError())
}
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{
_, 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)
}
+60 -22
View File
@@ -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)
}
+14 -1
View File
@@ -611,12 +611,25 @@ func (s *controlBatchSender) controlBatchWaitContext(requests []controlBatchRequ
base = s.stopCtx
}
ctx, cancel := context.WithCancel(base)
var remaining atomic.Int32
for _, item := range requests {
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, cancel))
stops = append(stops, context.AfterFunc(item.ctx, cancelWhenAllDone))
}
}
return ctx, func() {
for _, stop := range stops {
+7 -1
View File
@@ -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)
+132 -5
View File
@@ -8,31 +8,108 @@ 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
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()
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{
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
}
if size < 0 {
size = 0
}
for {
d.mu.Lock()
if d.closed {
d.mu.Unlock()
@@ -43,7 +120,12 @@ func (d *inboundDispatcher) Dispatch(source string, fn func()) bool {
worker = &inboundDispatchWorker{}
d.workers[source] = worker
}
worker.queue = append(worker.queue, fn)
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
@@ -54,6 +136,41 @@ func (d *inboundDispatcher) Dispatch(source string, fn func()) bool {
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
}
if d.sourceItems > 0 && worker.queued >= d.sourceItems {
return false
}
// 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
}
return true
}
func (d *inboundDispatcher) signalRoomLocked() {
select {
case d.roomCh <- struct{}{}:
default:
}
}
func (d *inboundDispatcher) run(source string, worker *inboundDispatchWorker) {
defer d.wg.Done()
@@ -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()
if !d.closed {
d.closed = true
close(d.closeCh)
}
d.signalRoomLocked()
d.mu.Unlock()
d.wg.Wait()
}
+231
View File
@@ -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")
}
}
+13 -2
View File
@@ -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,
}
+11
View File
@@ -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 {
+317
View File
@@ -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
}
+40
View File
@@ -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)
}
}
}
+52 -19
View File
@@ -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{
+184
View File
@@ -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
}
+17 -5
View File
@@ -3,6 +3,8 @@ package notify
const (
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"
)
@@ -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
return metadata, StreamMetadata{
recordStreamMetadataUseBatchAckKey: recordStreamMetadataEnabledValue,
response[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
}
+177
View File
@@ -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))
})
}
}
+318
View File
@@ -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)
}
}
+373
View File
@@ -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)
}
}
+23
View File
@@ -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)
}
+148 -144
View File
@@ -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
}
@@ -129,12 +135,20 @@ type recordStream struct {
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
@@ -161,6 +175,7 @@ type recordStream struct {
maxPendingApply int
remoteClosed bool
inboundClosed bool
readErr error
terminalErr error
}
@@ -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,
}
}
@@ -237,11 +260,14 @@ func WrapStreamAsRecord(stream Stream, opt RecordOpenOptions) (RecordStream, err
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,24 +354,14 @@ 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:
// 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
}
}
}
func (r *recordStream) Flush(ctx context.Context) error {
if r == 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,24 +719,25 @@ 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 {
if _, err := appendBatch(req); err != nil {
return err
}
}
}
}
for {
select {
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
+91 -3
View File
@@ -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)
+197
View File
@@ -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)
}
}
+3
View File
@@ -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()
+191 -82
View File
@@ -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()
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{
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()
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{
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)
+10 -1
View File
@@ -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) {
+3 -4
View File
@@ -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 {
+5 -5
View File
@@ -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
+17
View File
@@ -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)
}
+56 -52
View File
@@ -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
}
}
@@ -129,15 +122,26 @@ func serverStreamResetSender(s *ServerCommon, logical *LogicalConn, transport *T
return func(ctx context.Context, stream *streamHandle, message string) error {
req := StreamResetRequest{
StreamID: stream.ID(),
DataID: stream.dataIDSnapshot(),
Error: message,
RecordFailure: stream.recordResetFailure(),
}
if logical != nil {
_, err := sendStreamResetServerLogical(ctx, s, logical, req)
return err
}
if transport != nil {
_, err := sendStreamResetServerTransport(ctx, s, transport, req)
return err
}
_, 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 {
+52 -15
View File
@@ -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 {
+34 -1
View File
@@ -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)
+6 -2
View File
@@ -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,
+16 -4
View File
@@ -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)
+131 -20
View File
@@ -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,21 +873,29 @@ 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
if !s.applyResetState(resetErr) {
return s.resetErrSnapshot()
}
defer cancel()
if sendErr := resetFn(ctx, s, streamResetMessage(resetErr)); sendErr != nil {
return sendErr
}
}
s.markReset(resetErr)
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()
}
}
func (s *streamHandle) markRemoteClosed() {
if s == nil {
@@ -826,17 +927,25 @@ func (s *streamHandle) markPeerClosed() {
}
func (s *streamHandle) markReset(err error) {
if s == nil {
return
_ = s.applyResetState(err)
}
func (s *streamHandle) applyResetState(err error) bool {
if s == nil {
return false
}
s.acceptState.CompareAndSwap(0, 2)
s.mu.Lock()
if s.resetErr == nil {
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
}
s.finalizeOnce.Do(func() {
if s.cancel != nil {
s.cancel()
}
if s.runtime != nil {
s.runtime.remove(s.runtimeScope, s.id)
s.runtime.remove(s.runtimeScope, s)
}
})
}
func (s *streamHandle) waitReadable(ctx context.Context, notify <-chan struct{}, deadlineNotify <-chan struct{}, deadline time.Time) error {
+81 -14
View File
@@ -47,6 +47,9 @@ type streamBatchSender struct {
doneCh chan struct{}
stopOnce sync.Once
admissionMu sync.Mutex
admitting sync.WaitGroup
admissionClosed bool
flushMu sync.Mutex
queued atomic.Int64
errMu sync.Mutex
@@ -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
+140 -24
View File
@@ -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
}
@@ -41,6 +43,7 @@ type StreamResetRequest struct {
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
+40 -10
View File
@@ -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,
+19 -5
View File
@@ -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))
+64
View File
@@ -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)
}
}
+415
View File
@@ -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)
}
}
+265 -6
View File
@@ -11,22 +11,36 @@ import (
type streamRuntime struct {
rolePrefix string
seq atomic.Uint64
dataSeq 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
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 {
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
+38 -10
View File
@@ -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,
+6 -2
View File
@@ -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)
+6
View File
@@ -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
+4
View File
@@ -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 {
+12
View File
@@ -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
}
+39 -2
View File
@@ -12,6 +12,11 @@ type TransportConn struct {
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 ""
+43
View File
@@ -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()
+16
View File
@@ -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 {
+145
View File
@@ -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)