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:
@@ -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)
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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{
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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 ""
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user