fix(notify): 根治低带宽控制面阻塞与传输写入卡死

- 为控制消息增加优先级、公平调度、队列字节预算和自适应批处理
- 支持可取消的写门等待,收紧 shared/dedicated bulk、stream 和 Reply 写入边界
- 修复 bulk reset/close、连接 handoff 和安全 profile 切换时序
- 保留旧取消与超时错误契约,新增阶段化 TransportSendError
- 增加 ReplyCtx、ReplyObjCtx、写超时配置及黑洞连接和竞态回归测试
This commit is contained in:
2026-08-14 10:16:37 +08:00
parent 98ef9e7fcc
commit 0826e17063
45 changed files with 4231 additions and 293 deletions
+99 -21
View File
@@ -11,6 +11,10 @@ import (
"time"
)
// A zero public WriteTimeout preserves the legacy option contract, but a
// physical bulk write still needs a finite bound so a black-hole peer cannot
// hold the async worker (and CloseWrite) forever. The default is constant so
// concurrent bulk operations cannot observe a process-wide timeout change.
const (
BulkOpenSignalKey = "notify.bulk.open"
BulkCloseSignalKey = "notify.bulk.close"
@@ -26,6 +30,8 @@ const (
defaultBulkControlReadTimeout = 0
defaultBulkControlWriteTimeout = 0
defaultBulkAcceptReadyTimeout = 10 * time.Second
defaultBulkResetNotifyTimeout = 30 * time.Second
defaultBulkDataWriteTimeout = 2 * time.Minute
)
type BulkMetadata map[string]string
@@ -1278,9 +1284,13 @@ func (b *bulkHandle) close(full bool) error {
return nil
}
closeFn := b.closeFn
writeTimeout := b.writeTimeout
b.mu.Unlock()
if closeFn != nil && !b.dedicatedWriteHalfClosedSnapshot() {
if err := closeFn(context.Background(), b, true); err != nil && !errors.Is(err, errBulkNotFound) && !b.canIgnoreDedicatedCloseSendError(err) {
closeCtx, cancel := b.closeContext(writeTimeout)
err := closeFn(closeCtx, b, true)
cancel()
if err != nil && !errors.Is(err, errBulkNotFound) && !b.canIgnoreDedicatedCloseSendError(err) {
return err
}
}
@@ -1302,12 +1312,19 @@ func (b *bulkHandle) close(full bool) error {
return nil
}
closeFn := b.closeFn
writeTimeout := b.writeTimeout
b.mu.Unlock()
if err := b.waitPendingAsyncWrites(context.Background()); err != nil {
drainCtx, drainCancel := b.closeContext(writeTimeout)
err := b.waitPendingAsyncWrites(drainCtx)
drainCancel()
if err != nil {
return err
}
if closeFn != nil {
if err := closeFn(context.Background(), b, full); err != nil && !errors.Is(err, errBulkNotFound) && !b.canIgnoreDedicatedCloseSendError(err) {
closeCtx, cancel := b.closeContext(writeTimeout)
err := closeFn(closeCtx, b, full)
cancel()
if err != nil && !errors.Is(err, errBulkNotFound) && !b.canIgnoreDedicatedCloseSendError(err) {
return err
}
}
@@ -1348,13 +1365,29 @@ func (b *bulkHandle) Reset(err error) error {
}
resetFn := b.resetFn
b.mu.Unlock()
if resetFn != nil {
if sendErr := resetFn(context.Background(), b, bulkResetMessage(resetErr)); sendErr != nil {
return sendErr
}
if !b.applyResetState(resetErr) {
return b.resetErrSnapshot()
}
defer b.finalize()
if resetFn == nil {
return nil
}
timeout := defaultBulkResetNotifyTimeout
if b.writeTimeout > 0 && b.writeTimeout < timeout {
timeout = b.writeTimeout
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
done := make(chan error, 1)
go func() {
done <- resetFn(ctx, b, bulkResetMessage(resetErr))
}()
select {
case sendErr := <-done:
return sendErr
case <-ctx.Done():
return ctx.Err()
}
b.markReset(resetErr)
return nil
}
func (b *bulkHandle) Snapshot() BulkSnapshot {
@@ -1396,18 +1429,34 @@ func (b *bulkHandle) markReset(err error) {
if b == nil {
return
}
resetErr := bulkResetError(err)
b.mu.Lock()
if b.resetErr == nil {
b.resetErr = resetErr
b.clearBufferedDataLocked()
b.closeWriteStateLocked()
b.applyResetState(bulkResetError(err))
b.finalize()
}
func (b *bulkHandle) applyResetState(resetErr error) bool {
if b == nil {
return false
}
b.mu.Lock()
if b.resetErr != nil {
b.notifyFlowLocked()
b.mu.Unlock()
return false
}
b.resetErr = bulkResetError(resetErr)
b.clearBufferedDataLocked()
b.closeWriteStateLocked()
b.notifyFlowLocked()
b.mu.Unlock()
b.markAcceptReady(resetErr)
b.markAcceptReady(b.resetErrSnapshot())
b.notifyReadable()
b.finalize()
if b.cancel != nil {
b.cancel()
}
if b.writeCtxCancel != nil {
b.writeCtxCancel()
}
return true
}
func (b *bulkHandle) pushChunk(chunk []byte) error {
@@ -1818,7 +1867,12 @@ func (b *bulkHandle) finalize() {
if b.writeCtxCancel != nil {
b.writeCtxCancel()
}
if sender := b.clearDedicatedSender(); sender != nil {
sender := b.clearDedicatedSender()
conn, owned := b.clearDedicatedConn()
if conn != nil && owned {
_ = conn.Close()
}
if sender != nil {
sender.stop()
}
if b.client != nil && b.releaseDedicatedActiveReserved() {
@@ -1827,9 +1881,6 @@ func (b *bulkHandle) finalize() {
if b.client != nil {
b.client.releaseBulkDedicatedLane(b.dedicatedLaneIDSnapshot())
}
if conn, owned := b.clearDedicatedConn(); conn != nil && owned {
_ = conn.Close()
}
if b.runtime != nil {
b.runtime.remove(b.runtimeScope, b.id)
}
@@ -2296,6 +2347,33 @@ func bulkWriteContext(parent context.Context, timeout time.Duration) (context.Co
return ctx, cancel, nil
}
func bulkCloseContext(parent context.Context, timeout time.Duration) (context.Context, context.CancelFunc) {
if parent == nil {
parent = context.Background()
}
if timeout <= 0 {
return context.WithTimeout(parent, defaultBulkDataWriteTimeout)
}
return context.WithTimeout(parent, timeout)
}
func (b *bulkHandle) closeContext(timeout time.Duration) (context.Context, context.CancelFunc) {
if b == nil {
return bulkCloseContext(nil, timeout)
}
parent := b.Context()
b.mu.Lock()
gracefulPeerClose := b.remoteClosed || b.peerReadClosed
b.mu.Unlock()
if gracefulPeerClose {
// A peer half/full close cancels the bulk context as part of normal EOF
// delivery. Keep the final close notification alive so the peer can
// complete the protocol handshake, while retaining its write timeout.
parent = context.Background()
}
return bulkCloseContext(parent, timeout)
}
func normalizeBulkOpenRequest(req BulkOpenRequest) BulkOpenRequest {
req.Range = normalizeBulkRange(req.Range)
req.Metadata = cloneBulkMetadata(req.Metadata)
+22 -1
View File
@@ -418,6 +418,11 @@ func (s *bulkBatchSender) flush(requests []bulkBatchRequest) error {
}
}()
writeTimeout := s.transportWriteTimeout()
if writeTimeout <= 0 {
writeTimeout = defaultBulkDataWriteTimeout
}
requestDeadline := bulkBatchRequestsEarliestDeadline(requests)
writeDeadline := earlierWriteDeadline(writeDeadlineFromTimeout(writeTimeout), requestDeadline)
frames := make([][]byte, 0, len(payloads))
payloadBytes := 0
for _, payload := range payloads {
@@ -425,13 +430,25 @@ func (s *bulkBatchSender) flush(requests []bulkBatchRequest) error {
payloadBytes += len(payload.payload)
}
started := time.Now()
err = s.binding.withConnWriteLockDeadline(writeDeadlineFromTimeout(writeTimeout), func(conn net.Conn) error {
lockAcquired, err := s.binding.withConnWriteLockContextStopDeadlineManaged(context.Background(), s.stopCh, writeDeadline, func(conn net.Conn) error {
return writeFramedPayloadBatchUnlocked(conn, queue, frames)
})
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.
s.binding.closeConn()
}
return err
}
func bulkBatchRequestsEarliestDeadline(requests []bulkBatchRequest) time.Time {
var deadline time.Time
for _, req := range requests {
deadline = earlierWriteDeadline(deadline, req.deadline)
}
return deadline
}
func (s *bulkBatchSender) encodeRequests(requests []bulkBatchRequest) ([]bulkBatchEncodedPayload, error) {
if len(requests) == 0 {
return nil, nil
@@ -609,6 +626,10 @@ func (s *bulkBatchSender) stop() {
close(s.stopCh)
})
<-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.
s.flushMu.Lock()
s.flushMu.Unlock()
}
func (s *bulkBatchSender) failPending(err error) {
+53
View File
@@ -1,12 +1,27 @@
package notify
import (
"b612.me/stario"
"context"
"errors"
"math"
"net"
"sync"
"testing"
"time"
)
type bulkReleaseTrackingConn struct {
net.Conn
started chan struct{}
once sync.Once
}
func (c *bulkReleaseTrackingConn) Write(p []byte) (int, error) {
c.once.Do(func() { close(c.started) })
return c.Conn.Write(p)
}
func TestBulkOwnedChunkReleaseAfterRead(t *testing.T) {
bulk := newBulkHandle(context.Background(), newBulkRuntime("buffer-release-read"), clientFileScope(), BulkOpenRequest{
BulkID: "buffer-release-read",
@@ -117,3 +132,41 @@ func TestBulkReadDoesNotBlockOnAsyncWindowRelease(t *testing.T) {
}
}
}
func TestLegacyBulkReleaseHonorsBulkCancellation(t *testing.T) {
client := NewClient().(*ClientCommon)
if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil {
t.Fatal(err)
}
stopCtx, stopFn := context.WithCancel(context.Background())
defer stopFn()
pipeLeft, right := net.Pipe()
left := &bulkReleaseTrackingConn{Conn: pipeLeft, started: make(chan struct{})}
defer left.Close()
defer right.Close()
queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32)
client.setClientSessionRuntime(newClientSessionRuntime(left, stopCtx, stopFn, queue, 1))
client.markSessionStarted()
bulk := newBulkHandle(context.Background(), nil, "test", BulkOpenRequest{
BulkID: "legacy-release",
DataID: 1,
FastPathVersion: bulkFastPathVersionV1,
ChunkSize: 4,
WindowBytes: 4,
MaxInFlight: 1,
}, 0, nil, nil, 0, nil, nil, nil, nil, clientBulkReleaseSender(client))
bulk.maybeSendWindowRelease(4, true)
select {
case <-left.started:
case <-time.After(time.Second):
t.Fatal("legacy release did not enter physical write")
}
bulk.finalize()
select {
case <-bulk.releaseWorkerDone:
case <-time.After(150 * time.Millisecond):
t.Fatal("legacy bulk release worker ignored bulk cancellation while its send was blocked")
}
}
+33 -6
View File
@@ -793,31 +793,58 @@ func sendBulkResetServerTransport(ctx context.Context, s Server, transport *Tran
return decodeBulkResetResponse(msg)
}
func sendBulkReleaseClient(c Client, req BulkReleaseRequest) error {
func sendBulkReleaseClient(ctx context.Context, c *ClientCommon, req BulkReleaseRequest) error {
if c == nil {
return errBulkClientNil
}
return c.SendObj(BulkReleaseSignalKey, req)
data, err := encode(req)
if err != nil {
return err
}
_, err = c.sendWithContext(ctx, TransferMsg{
Key: BulkReleaseSignalKey,
Value: data,
Type: MSG_ASYNC,
})
return err
}
func sendBulkReleaseServerLogical(s Server, logical *LogicalConn, req BulkReleaseRequest) error {
func sendBulkReleaseServerLogical(ctx context.Context, s *ServerCommon, logical *LogicalConn, req BulkReleaseRequest) error {
if s == nil {
return errBulkServerNil
}
if logical == nil {
return errBulkLogicalConnNil
}
return s.SendObjLogical(logical, BulkReleaseSignalKey, req)
data, err := encode(req)
if err != nil {
return err
}
_, err = s.sendLogicalContext(ctx, logical, TransferMsg{
Key: BulkReleaseSignalKey,
Value: data,
Type: MSG_ASYNC,
}, 0)
return err
}
func sendBulkReleaseServerTransport(s Server, transport *TransportConn, req BulkReleaseRequest) error {
func sendBulkReleaseServerTransport(ctx context.Context, s *ServerCommon, transport *TransportConn, req BulkReleaseRequest) error {
if s == nil {
return errBulkServerNil
}
if transport == nil {
return errBulkTransportNil
}
return s.SendObjTransport(transport, BulkReleaseSignalKey, req)
data, err := encode(req)
if err != nil {
return err
}
_, err = s.sendTransportContext(ctx, transport, TransferMsg{
Key: BulkReleaseSignalKey,
Value: data,
Type: MSG_ASYNC,
})
return err
}
func decodeBulkOpenRequest(msg *Message) (BulkOpenRequest, error) {
+21 -2
View File
@@ -221,6 +221,9 @@ func writeBulkDedicatedRecordWithDeadline(conn net.Conn, payload []byte, deadlin
if conn == nil {
return net.ErrClosed
}
if deadline.IsZero() {
deadline = writeDeadlineFromTimeout(defaultBulkDataWriteTimeout)
}
return withRawConnWriteLockDeadline(conn, deadline, func(conn net.Conn) error {
var header [bulkDedicatedRecordHeaderLen]byte
copy(header[:4], bulkDedicatedRecordMagic)
@@ -648,6 +651,9 @@ func (c *ClientCommon) sendDedicatedBulkAttachRequest(ctx context.Context, conn
if bulk == nil {
return bulkAttachResponse{}, errBulkIDEmpty
}
if ctx == nil {
ctx = context.Background()
}
defer func() {
_ = conn.SetReadDeadline(time.Time{})
}()
@@ -670,7 +676,16 @@ func (c *ClientCommon) sendDedicatedBulkAttachRequest(ctx context.Context, conn
if err != nil {
return bulkAttachResponse{}, err
}
if err := writeFullToConn(conn, frame); err != nil {
if err := ctx.Err(); err != nil {
return bulkAttachResponse{}, err
}
deadline := earlierWriteDeadline(
writeDeadlineFromTimeout(c.maxWriteTimeoutSnapshot()),
contextDeadline(ctx),
)
if err := withRawConnWriteLockDeadline(conn, deadline, func(conn net.Conn) error {
return writeFullToConnUnlocked(conn, frame)
}); err != nil {
return bulkAttachResponse{}, err
}
if deadline, ok := ctx.Deadline(); ok {
@@ -1019,7 +1034,11 @@ func (s *ServerCommon) replyDedicatedBulkAttachDetached(client *LogicalConn, con
if err != nil {
return err
}
return withRawConnWriteLockDeadline(conn, writeDeadlineFromTimeout(client.maxWriteTimeoutSnapshot()), func(conn net.Conn) error {
deadline := earlierWriteDeadline(
writeDeadlineFromTimeout(defaultBulkDedicatedHelloTimeout),
writeDeadlineFromTimeout(client.maxWriteTimeoutSnapshot()),
)
return withRawConnWriteLockDeadline(conn, deadline, func(conn net.Conn) error {
return writeFullToConnUnlocked(conn, frame)
})
}
+3 -3
View File
@@ -81,12 +81,12 @@ func (s *bulkDedicatedSidecar) close() {
return
}
s.closeOnce.Do(func() {
if sender := s.laneSenderSnapshot(); sender != nil {
sender.stop()
}
if s.conn != nil {
_ = s.conn.Close()
}
if sender := s.laneSenderSnapshot(); sender != nil {
sender.stop()
}
})
}
+31 -2
View File
@@ -370,7 +370,7 @@ func (s *ServerCommon) sendFastBulkDataTransport(ctx context.Context, logical *L
if err != nil {
return err
}
return s.writeEnvelopePayload(logical, transport, nil, payload)
return s.writeEnvelopePayloadContext(ctx, logical, transport, nil, payload)
}
func (s *ServerCommon) sendFastBulkWriteTransport(ctx context.Context, logical *LogicalConn, transport *TransportConn, dataID uint64, startSeq uint64, chunkSize int, fastPathVersion uint8, payload []byte, payloadOwned bool) (int, error) {
@@ -439,7 +439,7 @@ func (s *ServerCommon) sendFastBulkControlTransport(ctx context.Context, logical
if err != nil {
return err
}
return s.writeEnvelopePayload(logical, transport, nil, encoded)
return s.writeEnvelopePayloadContext(ctx, logical, transport, nil, encoded)
}
func (s *ServerCommon) encodeBulkFastControlPayloadLogical(logical *LogicalConn, frameType uint8, flags uint8, dataID uint64, seq uint64, payload []byte) ([]byte, error) {
@@ -474,6 +474,9 @@ func transportFastPayloadMagic(payload []byte) string {
func (c *ClientCommon) decryptTransportPayloadPooled(payload []byte, release func()) ([]byte, func(), error) {
profile := c.clientTransportProtectionSnapshot()
if fallback := c.inboundTransitionProfile.Load(); fallback != nil {
return decryptTransportPayloadWithFallbackPooled(profile, *fallback, payload, release)
}
return decryptTransportPayloadCodecPooled(profile.mode, profile.runtime, profile.msgDe, profile.secretKey, payload, release)
}
@@ -484,9 +487,35 @@ func (s *ServerCommon) decryptTransportPayloadLogicalPooled(logical *LogicalConn
}
return nil, nil, errTransportDetached
}
if fallback := logical.inboundTransitionProfile.Load(); fallback != nil {
profile := logical.transportProtectionProfileSnapshot()
return decryptTransportPayloadWithFallbackPooled(profile, *fallback, payload, release)
}
return decryptTransportPayloadCodecPooled(logical.protectionModeSnapshot(), logical.modernPSKRuntimeSnapshot(), logical.msgDeSnapshot(), logical.secretKeySnapshot(), payload, release)
}
func decryptTransportPayloadWithFallbackPooled(primary transportProtectionProfile, fallback transportProtectionProfile, payload []byte, release func()) ([]byte, func(), error) {
profiles := [...]transportProtectionProfile{primary, fallback}
for _, profile := range profiles {
plain, plainRelease, err := decryptTransportPayloadCodecOwnedPooled(profile.mode, profile.runtime, profile.msgDe, profile.secretKey, payload)
if err != nil {
continue
}
if profile.mode == ProtectionExternal {
plain = append([]byte(nil), plain...)
plainRelease = nil
}
if release != nil {
release()
}
return plain, plainRelease, nil
}
if release != nil {
release()
}
return nil, nil, errTransportPayloadDecryptFailed
}
func (c *ClientCommon) tryDispatchBorrowedTransportPlain(plain []byte, release func()) bool {
switch transportFastPayloadMagic(plain) {
case bulkFastPayloadMagic, bulkFastBatchMagic:
+199 -1
View File
@@ -6,6 +6,7 @@ import (
"io"
"net"
"strings"
"sync"
"testing"
"time"
)
@@ -818,6 +819,61 @@ func TestBulkWritePrefersResetErrorOverContextCanceled(t *testing.T) {
}
}
func TestBulkResetWakesLocalStateBeforeRemoteNotificationCompletes(t *testing.T) {
wantErr := errors.New("local reset must win")
resetStarted := make(chan struct{})
resetUnblock := make(chan struct{})
var unblockOnce sync.Once
unblock := func() { unblockOnce.Do(func() { close(resetUnblock) }) }
defer unblock()
bulk := newBulkHandle(context.Background(), nil, "test", BulkOpenRequest{
BulkID: "bulk-reset-local-first",
DataID: 1,
Range: BulkRange{
Length: 1,
},
}, 0, nil, nil, 0, nil, func(context.Context, *bulkHandle, string) error {
close(resetStarted)
<-resetUnblock
return nil
}, nil, nil, nil)
readDone := make(chan error, 1)
go func() {
_, err := bulk.Read(make([]byte, 1))
readDone <- err
}()
resetDone := make(chan error, 1)
go func() {
resetDone <- bulk.Reset(wantErr)
}()
select {
case <-resetStarted:
case <-time.After(time.Second):
t.Fatal("remote reset notification did not start")
}
select {
case <-bulk.Context().Done():
case <-time.After(100 * time.Millisecond):
t.Fatal("bulk context stayed live while remote reset notification was blocked")
}
select {
case err := <-readDone:
if !errors.Is(err, wantErr) {
t.Fatalf("bulk Read error=%v, want %v", err, wantErr)
}
case <-time.After(100 * time.Millisecond):
t.Fatal("bulk Read stayed blocked while remote reset notification was blocked")
}
unblock()
if err := <-resetDone; err != nil {
t.Fatalf("Reset returned error: %v", err)
}
}
func TestDedicatedBulkWaitReadyPrefersClosedPipeOverContextCanceled(t *testing.T) {
bulk := newBulkHandle(context.Background(), nil, "test", BulkOpenRequest{
BulkID: "bulk-dedicated-ready-close",
@@ -892,6 +948,111 @@ func TestBulkReadWaitingLocalClosePrefersClosedPipeOverContextCanceled(t *testin
}
}
func TestBulkDefaultWriteContextDoesNotAddTimer(t *testing.T) {
parent := context.Background()
ctx, cancel, err := bulkWriteContext(parent, 0)
if err != nil {
t.Fatalf("bulkWriteContext returned error: %v", err)
}
defer cancel()
if ctx != parent {
t.Fatal("zero WriteTimeout wrapped the parent context")
}
if _, ok := ctx.Deadline(); ok {
t.Fatal("zero WriteTimeout added a per-write deadline")
}
}
func TestBulkDefaultCloseContextIsBounded(t *testing.T) {
ctx, cancel := bulkCloseContext(context.Background(), 0)
defer cancel()
deadline, ok := ctx.Deadline()
if !ok {
t.Fatal("default bulk close context has no deadline")
}
remaining := time.Until(deadline)
if remaining < defaultBulkDataWriteTimeout-time.Second || remaining > defaultBulkDataWriteTimeout+time.Second {
t.Fatalf("default bulk close deadline remaining=%v, want about %v", remaining, defaultBulkDataWriteTimeout)
}
}
func TestBulkCloseWriteHonorsBulkContextCancellation(t *testing.T) {
parent, cancelParent := context.WithCancel(context.Background())
defer cancelParent()
closeStarted := make(chan struct{})
bulk := newBulkHandle(
parent,
nil,
clientFileScope(),
BulkOpenRequest{BulkID: "close-write-context-bound", DataID: 1},
0,
nil,
nil,
0,
func(ctx context.Context, _ *bulkHandle, _ bool) error {
close(closeStarted)
<-ctx.Done()
return ctx.Err()
},
nil,
nil,
nil,
nil,
)
done := make(chan error, 1)
go func() { done <- bulk.CloseWrite() }()
select {
case <-closeStarted:
case <-time.After(time.Second):
t.Fatal("bulk close notification did not start")
}
cancelParent()
select {
case err := <-done:
if !errors.Is(err, context.Canceled) {
t.Fatalf("Bulk.CloseWrite error=%v, want context canceled", err)
}
case <-time.After(time.Second):
t.Fatal("Bulk.CloseWrite ignored the canceled bulk context")
}
}
func TestBulkCloseAfterPeerEOFUsesBoundedCleanupContext(t *testing.T) {
closeCalled := false
bulk := newBulkHandle(
context.Background(),
nil,
clientFileScope(),
BulkOpenRequest{BulkID: "close-after-peer-eof", DataID: 1},
0,
nil,
nil,
0,
func(ctx context.Context, _ *bulkHandle, full bool) error {
if !full {
return errors.New("expected full close after peer EOF")
}
if err := ctx.Err(); err != nil {
return err
}
closeCalled = true
return nil
},
nil,
nil,
nil,
nil,
)
bulk.markPeerClosed()
if err := bulk.Close(); err != nil {
t.Fatalf("Bulk.Close after peer EOF failed: %v", err)
}
if !closeCalled {
t.Fatal("Bulk.Close skipped the final peer notification after EOF")
}
}
func TestBulkReleaseControlRoundTripTransport(t *testing.T) {
server := NewServer().(*ServerCommon)
if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil {
@@ -951,7 +1112,7 @@ func TestBulkReleaseControlRoundTripTransport(t *testing.T) {
clientHandle.outboundInFlight = 1
clientHandle.mu.Unlock()
if err := sendBulkReleaseServerTransport(server, accepted.TransportConn, BulkReleaseRequest{
if err := sendBulkReleaseServerTransport(context.Background(), server, accepted.TransportConn, BulkReleaseRequest{
BulkID: serverHandle.ID(),
DataID: serverHandle.dataIDSnapshot(),
Bytes: chunkSize,
@@ -1549,6 +1710,43 @@ func TestBulkDedicatedClientFullCloseAfterCloseWriteDoesNotResetTCP(t *testing.T
}
}
func TestBulkCloseWriteUsesWriteTimeoutForControlNotification(t *testing.T) {
const timeout = 40 * time.Millisecond
bulk := newBulkHandle(
context.Background(),
nil,
clientFileScope(),
BulkOpenRequest{
BulkID: "close-write-timeout",
WriteTimeout: timeout,
Range: BulkRange{Length: 1},
},
0,
nil,
nil,
0,
func(ctx context.Context, _ *bulkHandle, _ bool) error {
<-ctx.Done()
return ctx.Err()
},
nil,
nil,
nil,
nil,
)
done := make(chan error, 1)
go func() { done <- bulk.CloseWrite() }()
select {
case err := <-done:
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("Bulk.CloseWrite error=%v, want deadline exceeded", err)
}
case <-time.After(time.Second):
t.Fatal("Bulk.CloseWrite remained blocked past WriteTimeout")
}
}
func TestBulkSharedConcurrentWritersWithSlowReceiver(t *testing.T) {
server := NewServer().(*ServerCommon)
if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil {
+2
View File
@@ -28,6 +28,7 @@ type ClientCommon struct {
maxReadTimeout time.Duration
maxWriteTimeout time.Duration
keyExchangeFn func(c Client) error
linkMu sync.RWMutex
linkFns map[string]func(message *Message)
defaultFns func(message *Message)
msgEn func([]byte, []byte) []byte
@@ -39,6 +40,7 @@ type ClientCommon struct {
handshakeRsaPubKey []byte
SecretKey []byte
transportProtection atomic.Pointer[transportProtectionProfile]
inboundTransitionProfile atomic.Pointer[transportProtectionProfile]
peerAttachSecurity atomic.Pointer[peerAttachSecurityState]
securityBootstrap transportProtectionProfile
securitySteady transportProtectionProfile
+1 -1
View File
@@ -333,7 +333,7 @@ func clientBulkReleaseSender(c *ClientCommon) bulkReleaseSender {
}
return c.sendFastBulkControl(ctx, bulkFastPayloadTypeRelease, 0, bulk.dataIDSnapshot(), 0, bulk.fastPathVersionSnapshot(), payload)
}
return sendBulkReleaseClient(c, BulkReleaseRequest{
return sendBulkReleaseClient(ctx, c, BulkReleaseRequest{
BulkID: bulk.ID(),
DataID: bulk.dataIDSnapshot(),
Bytes: bytes,
+4 -2
View File
@@ -32,12 +32,14 @@ func (c *ClientCommon) ShowError(std bool) {
}
func (c *ClientCommon) SetDefaultLink(fn func(message *Message)) {
c.linkMu.Lock()
defer c.linkMu.Unlock()
c.defaultFns = fn
}
func (c *ClientCommon) SetLink(key string, fn func(*Message)) {
c.mu.Lock()
defer c.mu.Unlock()
c.linkMu.Lock()
defer c.linkMu.Unlock()
c.linkFns[key] = fn
}
+6 -1
View File
@@ -107,6 +107,11 @@ func (c *LogicalConn) detachServerOwnedTransport() {
if c == nil {
return
}
c.closeTransport()
conn := c.transportSnapshot()
// Revoke read-loop ownership before closing the socket. Otherwise the close
// error can race with detach and incorrectly stop the logical session.
c.clearSessionRuntimeTransport()
if conn != nil {
_ = conn.Close()
}
}
+6 -3
View File
@@ -24,12 +24,15 @@ func (c *ClientCommon) dispatchMsg(message Message) {
callFn := func(fn func(*Message)) {
fn(&message)
}
c.linkMu.RLock()
fn, ok := c.linkFns[message.Key]
if ok {
defaultFn := c.defaultFns
c.linkMu.RUnlock()
if ok && fn != nil {
callFn(fn)
}
if c.defaultFns != nil {
callFn(c.defaultFns)
if defaultFn != nil {
callFn(defaultFn)
}
}
+10 -1
View File
@@ -289,7 +289,7 @@ func (c *ClientCommon) startClientWithConnSource(conn net.Conn, source *clientCo
queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32)
c.setClientConnectSource(source)
rt := newClientSessionRuntime(conn, stopCtx, stopFn, queue, epoch)
c.setClientSessionRuntime(rt)
c.setClientSessionRuntimeWithCloseOld(rt, true)
c.resetClientStopState()
c.markSessionStarted()
return c.clientPostInit(rt)
@@ -360,7 +360,16 @@ func (c *ClientCommon) bootstrapClientTransportRuntime(rt *clientSessionRuntime,
if err := c.announceClientPeerIdentity(); err != nil {
return c.failClientTransportBootstrap(rt, stopSessionOnFailure, "peer attach failed", err)
}
var transitionProfile *transportProtectionProfile
if c.securityConfigured {
transitionProfile = c.installInboundTransitionProfile(c.clientTransportProtectionSnapshot())
}
c.activateClientSteadyTransportProtection()
if transitionProfile != nil {
time.AfterFunc(peerAttachTransitionFallbackTTL, func() {
c.clearInboundTransitionProfile(transitionProfile)
})
}
return nil
}
+30 -11
View File
@@ -9,9 +9,20 @@ import (
)
func (c *ClientCommon) send(msg TransferMsg) (WaitMsg, error) {
return c.sendWithContext(context.Background(), msg)
}
func (c *ClientCommon) sendWithContext(ctx context.Context, msg TransferMsg) (WaitMsg, error) {
return c.sendWithContextTimeout(ctx, msg, 0)
}
func (c *ClientCommon) sendWithContextTimeout(ctx context.Context, msg TransferMsg, writeTimeout time.Duration) (WaitMsg, error) {
if err := c.ensureClientSendReady(); err != nil {
return WaitMsg{}, err
}
if ctx == nil {
ctx = context.Background()
}
var wait WaitMsg
if msg.Type != MSG_SYNC_REPLY && msg.Type != MSG_KEY_CHANGE && msg.Type != MSG_SYS_REPLY || msg.ID == 0 {
msg.ID = atomic.AddUint64(&c.msgID, 1)
@@ -20,6 +31,8 @@ func (c *ClientCommon) send(msg TransferMsg) (WaitMsg, error) {
if err != nil {
return WaitMsg{}, err
}
env.controlCtx = ctx
env.controlTimeout = writeTimeout
if requiresSignalReplyWait(msg) {
wait = c.getPendingWaitPool().createAndStore(msg)
}
@@ -42,9 +55,9 @@ func (c *ClientCommon) sendEnvelope(env Envelope) error {
return err
}
if batchedControlEnvelope(env) {
return c.writeControlPayloadToTransport(payload)
return c.writeControlPayloadToTransportTimeout(env.controlContext(), payload, env.controlPriority, env.controlTimeout)
}
return c.writePayloadToTransport(payload)
return c.writePayloadToTransportContextTimeout(env.controlContext(), payload, env.controlTimeout)
}
func (c *ClientCommon) dispatchEnvelope(env Envelope, now time.Time) {
@@ -90,12 +103,18 @@ func (c *ClientCommon) Send(key string, value MsgVal) error {
}
func (c *ClientCommon) sendWait(msg TransferMsg, timeout time.Duration) (Message, error) {
data, err := c.send(msg)
ctx := context.Background()
cancel := func() {}
if timeout != 0 {
ctx, cancel = context.WithTimeout(ctx, timeout)
}
defer cancel()
data, err := c.sendWithContext(ctx, msg)
if err != nil {
return Message{}, err
return Message{}, publicContextSendError(ctx, err)
}
stopCh := sessionStopChan(c.clientStopContextSnapshot())
if timeout.Seconds() == 0 {
if timeout == 0 {
msg, ok := <-data.Reply
if !ok {
return msg, pendingWaitClosedErrorWith(stopCh, clientTransportDetachedError(c))
@@ -103,7 +122,7 @@ func (c *ClientCommon) sendWait(msg TransferMsg, timeout time.Duration) (Message
return msg, nil
}
select {
case <-time.After(timeout):
case <-ctx.Done():
c.getPendingWaitPool().removeAndClose(data.TransferMsg.ID)
return Message{}, os.ErrDeadlineExceeded
case <-stopCh:
@@ -117,14 +136,14 @@ func (c *ClientCommon) sendWait(msg TransferMsg, timeout time.Duration) (Message
}
func (c *ClientCommon) sendCtx(msg TransferMsg, ctx context.Context) (Message, error) {
data, err := c.send(msg)
if err != nil {
return Message{}, err
}
stopCh := sessionStopChan(c.clientStopContextSnapshot())
if ctx == nil {
ctx = context.Background()
}
data, err := c.sendWithContext(ctx, msg)
if err != nil {
return Message{}, publicContextSendError(ctx, err)
}
stopCh := sessionStopChan(c.clientStopContextSnapshot())
select {
case <-ctx.Done():
c.getPendingWaitPool().removeAndClose(data.TransferMsg.ID)
+7 -3
View File
@@ -50,6 +50,10 @@ func prepareClientSessionRuntime(rt *clientSessionRuntime) *clientSessionRuntime
}
func (c *ClientCommon) setClientSessionRuntime(rt *clientSessionRuntime) {
c.setClientSessionRuntimeWithCloseOld(rt, false)
}
func (c *ClientCommon) setClientSessionRuntimeWithCloseOld(rt *clientSessionRuntime, closeOld bool) {
if c == nil || rt == nil {
return
}
@@ -69,7 +73,7 @@ func (c *ClientCommon) setClientSessionRuntime(rt *clientSessionRuntime) {
c.conn = rt.conn
}
if oldBinding != nil {
oldBinding.stopBackgroundWorkers()
stopReplacedTransportBinding(oldBinding, rt.transport, closeOld)
}
}
@@ -155,7 +159,7 @@ func (c *ClientCommon) clearClientSessionRuntimeTransport() {
next.conn = nil
next.transportStopCtx = nil
next.transportStopFn = nil
c.setClientSessionRuntime(&next)
c.setClientSessionRuntimeWithCloseOld(&next, true)
}
func (c *ClientCommon) clearClientSessionRuntimeQueue() {
@@ -199,7 +203,7 @@ func (c *ClientCommon) attachClientSessionTransport(conn net.Conn) error {
next.transportStopCtx = nil
next.transportStopFn = nil
next.suppressGoodByeOnStop = &atomic.Bool{}
c.setClientSessionRuntime(&next)
c.setClientSessionRuntimeWithCloseOld(&next, true)
if oldConn := oldBinding.connSnapshot(); oldConn != nil && oldConn != conn {
_ = oldConn.Close()
}
+95 -2
View File
@@ -6,11 +6,47 @@ import (
"io"
"math"
"net"
"sync"
"sync/atomic"
"testing"
"time"
)
type reattachBlockingConn struct {
started chan struct{}
finished chan struct{}
once sync.Once
}
func newReattachBlockingConn() *reattachBlockingConn {
return &reattachBlockingConn{started: make(chan struct{}), finished: make(chan struct{})}
}
func (c *reattachBlockingConn) Read([]byte) (int, error) { return 0, net.ErrClosed }
func (c *reattachBlockingConn) Close() error {
if c != nil {
c.once.Do(func() { close(c.finished) })
}
return nil
}
func (c *reattachBlockingConn) LocalAddr() net.Addr { return nil }
func (c *reattachBlockingConn) RemoteAddr() net.Addr { return nil }
func (c *reattachBlockingConn) SetDeadline(time.Time) error { return nil }
func (c *reattachBlockingConn) SetReadDeadline(time.Time) error { return nil }
func (c *reattachBlockingConn) SetWriteDeadline(time.Time) error { return nil }
func (c *reattachBlockingConn) Write([]byte) (int, error) {
closeOnce := func() {
select {
case <-c.started:
default:
close(c.started)
}
}
closeOnce()
<-c.finished
return 0, net.ErrClosed
}
func TestClientWriteToTransportUsesRuntimeConn(t *testing.T) {
client := NewClient().(*ClientCommon)
fallbackLeft, fallbackRight := net.Pipe()
@@ -23,12 +59,12 @@ func TestClientWriteToTransportUsesRuntimeConn(t *testing.T) {
client.conn = fallbackLeft
runtimeCtx, runtimeCancel := context.WithCancel(context.Background())
defer runtimeCancel()
client.setClientSessionRuntime(&clientSessionRuntime{
client.setClientSessionRuntimeWithCloseOld(&clientSessionRuntime{
conn: runtimeLeft,
stopCtx: runtimeCtx,
stopFn: runtimeCancel,
epoch: 1,
})
}, true)
payload := []byte("runtime-conn")
recvCh := make(chan []byte, 1)
@@ -350,3 +386,60 @@ func TestSetClientSessionRuntimeStopsOldBindingWorkersOnReattach(t *testing.T) {
t.Fatalf("old sender submit after reattach = %v, want %v", err, errTransportDetached)
}
}
func TestSetClientSessionRuntimeClosesOldConnBeforeStoppingWorkers(t *testing.T) {
client := NewClient().(*ClientCommon)
stopCtx, stopFn := context.WithCancel(context.Background())
defer stopFn()
queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32)
oldConn := newReattachBlockingConn()
oldBinding := newTransportBinding(oldConn, queue)
oldSender := newTestBulkBatchSender(oldBinding)
oldBinding.bulkMu.Lock()
oldBinding.bulkSender = oldSender
oldBinding.bulkMu.Unlock()
client.setClientSessionRuntime(&clientSessionRuntime{
transport: oldBinding,
conn: oldConn,
stopCtx: stopCtx,
stopFn: stopFn,
queue: queue,
epoch: 1,
})
writeDone := make(chan error, 1)
go func() {
writeDone <- oldSender.submitData(context.Background(), 1, 1, bulkFastPathVersionV1, []byte("blocked"))
}()
select {
case <-oldConn.started:
case <-time.After(time.Second):
t.Fatal("old sender did not enter physical write")
}
newLeft, newRight := net.Pipe()
defer newLeft.Close()
defer newRight.Close()
newBinding := newTransportBinding(newLeft, queue)
reattachDone := make(chan struct{})
go func() {
client.setClientSessionRuntimeWithCloseOld(&clientSessionRuntime{
transport: newBinding,
conn: newLeft,
stopCtx: stopCtx,
stopFn: stopFn,
queue: queue,
epoch: 2,
}, true)
close(reattachDone)
}()
select {
case <-reattachDone:
case <-time.After(time.Second):
t.Fatal("reattach remained blocked while old sender was writing")
}
select {
case <-writeDone:
case <-time.After(time.Second):
t.Fatal("old sender did not exit after old connection close")
}
}
+24 -8
View File
@@ -239,6 +239,14 @@ func (c *ClientCommon) writeToTransport(data []byte) error {
}
func (c *ClientCommon) writePayloadToTransport(payload []byte) error {
return c.writePayloadToTransportContext(context.Background(), payload)
}
func (c *ClientCommon) writePayloadToTransportContext(ctx context.Context, payload []byte) error {
return c.writePayloadToTransportContextTimeout(ctx, payload, 0)
}
func (c *ClientCommon) writePayloadToTransportContextTimeout(ctx context.Context, payload []byte, writeTimeout time.Duration) error {
binding := c.clientTransportBindingSnapshot()
if binding == nil {
return net.ErrClosed
@@ -247,15 +255,23 @@ func (c *ClientCommon) writePayloadToTransport(payload []byte) error {
if queue == nil {
return errClientSessionQueueUnavailable
}
return binding.withConnWriteLock(func(conn net.Conn) error {
if c.maxWriteTimeout.Seconds() != 0 {
_ = conn.SetWriteDeadline(time.Now().Add(c.maxWriteTimeout))
}
writeTimeout = shorterPositiveDuration(c.maxWriteTimeoutSnapshot(), writeTimeout)
lockAcquired, err := binding.withConnWriteLockContextStopTimeout(ctx, nil, writeTimeout, func(conn net.Conn) error {
return writeFramedPayloadUnlocked(conn, queue, payload)
})
if lockAcquired && err != nil {
if conn := binding.connSnapshot(); conn != nil && !isPacketTransportConn(conn) {
binding.closeConn()
}
}
return err
}
func (c *ClientCommon) writeControlPayloadToTransport(payload []byte) error {
func (c *ClientCommon) writeControlPayloadToTransport(ctx context.Context, payload []byte, priority controlPriority) error {
return c.writeControlPayloadToTransportTimeout(ctx, payload, priority, 0)
}
func (c *ClientCommon) writeControlPayloadToTransportTimeout(ctx context.Context, payload []byte, priority controlPriority, writeTimeout time.Duration) error {
binding := c.clientTransportBindingSnapshot()
if binding == nil {
return net.ErrClosed
@@ -266,11 +282,11 @@ func (c *ClientCommon) writeControlPayloadToTransport(payload []byte) error {
}
conn := binding.connSnapshot()
if conn == nil || isPacketTransportConn(conn) {
return c.writePayloadToTransport(payload)
return c.writePayloadToTransportContextTimeout(ctx, payload, writeTimeout)
}
sender := binding.controlBatchSenderSnapshot()
if sender == nil {
return c.writePayloadToTransport(payload)
return c.writePayloadToTransportContextTimeout(ctx, payload, writeTimeout)
}
return sender.submit(payload, writeDeadlineFromTimeout(c.maxWriteTimeout))
return sender.submitContext(ctx, payload, shorterPositiveDuration(c.maxWriteTimeoutSnapshot(), writeTimeout), priority)
}
+733 -78
View File
@@ -1,152 +1,727 @@
package notify
import (
"bytes"
"context"
"fmt"
"net"
"sync"
"sync/atomic"
"time"
)
const controlBatchMaxPayloads = 16
const (
controlBatchMaxPayloads = 16
controlBatchMaxPayloadBytes = 32 * 1024 * 1024
controlBatchMaxQueuedBytes = 32 * 1024 * 1024
controlBatchCriticalReservedBytes = 1 * 1024 * 1024
controlBatchMaxCriticalBurst = 16
)
type controlPriority uint8
const (
controlPriorityNormal controlPriority = iota
controlPriorityCritical
)
const (
controlBatchRequestQueued int32 = iota
controlBatchRequestStarted
controlBatchRequestCanceled
)
type controlBatchRequestState struct {
value atomic.Int32
}
type controlBatchRequest struct {
payload []byte
deadline time.Time
done chan error
ctx context.Context
payload []byte
writeTime time.Duration
priority controlPriority
done chan error
state *controlBatchRequestState
queueSize int64
}
type controlBatchSender struct {
binding *transportBinding
reqCh chan controlBatchRequest
stopCh chan struct{}
doneCh chan struct{}
binding *transportBinding
normalCh chan controlBatchRequest
criticalCh chan controlBatchRequest
stopCh chan struct{}
doneCh chan struct{}
budgetWake chan struct{}
stopCtx context.Context
stopCancel context.CancelFunc
stopOnce sync.Once
errMu sync.Mutex
err error
stopOnce sync.Once
flushMu sync.Mutex
admissionMu sync.Mutex
admitting sync.WaitGroup
admissionClosed bool
queued atomic.Int64
queuedSize atomic.Int64
errMu sync.Mutex
err error
}
func newControlBatchSender(binding *transportBinding) *controlBatchSender {
stopCtx, stopCancel := context.WithCancel(context.Background())
sender := &controlBatchSender{
binding: binding,
reqCh: make(chan controlBatchRequest, controlBatchMaxPayloads*4),
stopCh: make(chan struct{}),
doneCh: make(chan struct{}),
binding: binding,
normalCh: make(chan controlBatchRequest, controlBatchMaxPayloads*4),
criticalCh: make(chan controlBatchRequest, controlBatchMaxPayloads*2),
stopCh: make(chan struct{}),
doneCh: make(chan struct{}),
budgetWake: make(chan struct{}, 1),
stopCtx: stopCtx,
stopCancel: stopCancel,
}
go sender.run()
return sender
}
func (s *controlBatchSender) submit(payload []byte, deadline time.Time) error {
func (s *controlBatchSender) submit(payload []byte, writeTimeout time.Duration) error {
return s.submitContext(context.Background(), payload, writeTimeout, controlPriorityNormal)
}
func (s *controlBatchSender) submitContext(ctx context.Context, payload []byte, writeTimeout time.Duration, priority controlPriority) error {
if s == nil {
return errTransportDetached
return newTransportSendError(TransportSendStageTransport, errTransportDetached)
}
req := controlBatchRequest{
payload: payload,
deadline: deadline,
done: make(chan error, 1),
if ctx == nil {
ctx = context.Background()
}
if err := s.errSnapshot(); err != nil {
return err
}
select {
case <-s.stopCh:
return s.stoppedErr()
case s.reqCh <- req:
if err := ctx.Err(); err != nil {
return newTransportSendError(TransportSendStageQueue, err)
}
return <-req.done
if len(payload) > controlBatchMaxPayloadBytes {
return newTransportSendError(
TransportSendStageQueue,
fmt.Errorf("control payload is %d bytes, maximum is %d; use stream or bulk transfer", len(payload), controlBatchMaxPayloadBytes),
)
}
req := controlBatchRequest{
ctx: ctx,
payload: payload,
writeTime: maxDuration(0, writeTimeout),
priority: priority,
queueSize: int64(maxInt(1, len(payload))),
}
if submitted, err := s.tryDirectSubmit(req); submitted {
return err
}
queueCtx, cancelQueue := contextWithTimeoutUpperBound(ctx, req.writeTime)
defer cancelQueue()
req.ctx = queueCtx
if req.ctx.Done() != nil {
req.payload = bytes.Clone(payload)
}
s.queued.Add(1)
if err := s.reserveQueueBytes(req.ctx, req.queueSize, priority); err != nil {
s.queued.Add(-1)
return err
}
req.done = make(chan error, 1)
req.state = &controlBatchRequestState{}
queued := false
defer func() {
if !queued {
s.queued.Add(-1)
s.releaseQueueBytes(req.queueSize)
}
}()
if !s.beginAdmission() {
return s.stoppedErr()
}
select {
case <-req.ctx.Done():
s.endAdmission()
return newTransportSendError(TransportSendStageQueue, req.ctx.Err())
case <-s.stopCh:
s.endAdmission()
return s.stoppedErr()
case s.requestChannel(priority) <- req:
queued = true
s.endAdmission()
}
select {
case err := <-req.done:
return err
case <-s.stopCh:
if req.tryCancel() {
return s.stoppedErr()
}
if req.ctx != nil && req.ctx.Done() != nil {
return s.stoppedErr()
}
return <-req.done
case <-req.ctx.Done():
if req.tryCancel() {
return newTransportSendError(TransportSendStageQueue, req.ctx.Err())
}
return newTransportSendError(TransportSendStageQueue, req.ctx.Err())
}
}
func contextWithTimeoutUpperBound(ctx context.Context, timeout time.Duration) (context.Context, context.CancelFunc) {
if ctx == nil {
ctx = context.Background()
}
if timeout <= 0 {
return ctx, func() {}
}
deadline := time.Now().Add(timeout)
if current, ok := ctx.Deadline(); ok && !deadline.Before(current) {
return ctx, func() {}
}
return context.WithDeadline(ctx, deadline)
}
func (s *controlBatchSender) tryDirectSubmit(req controlBatchRequest) (bool, error) {
if s == nil {
return true, newTransportSendError(TransportSendStageTransport, errTransportDetached)
}
if err := s.errSnapshot(); err != nil {
return true, err
}
select {
case <-req.ctx.Done():
return true, newTransportSendError(TransportSendStageQueue, req.ctx.Err())
case <-s.stopCh:
return true, s.stoppedErr()
default:
}
if req.ctx.Done() != nil {
return false, nil
}
if s.queued.Load() != 0 || !s.flushMu.TryLock() {
return false, nil
}
defer s.flushMu.Unlock()
if s.queued.Load() != 0 {
return false, nil
}
if err := s.errSnapshot(); err != nil {
return true, err
}
if err := s.flushDirect(req); err != nil {
if stage, ok := TransportSendErrorStage(err); ok && stage == TransportSendStageQueue {
return true, err
}
s.markFailed(err)
s.waitAdmissions()
s.failPending(err, nil, nil)
return true, err
}
return true, nil
}
func (s *controlBatchSender) run() {
defer close(s.doneCh)
var pendingNormal []controlBatchRequest
var pendingCritical []controlBatchRequest
criticalBurst := 0
for {
req, ok := s.nextRequest()
forceNormal := criticalBurst >= controlBatchMaxCriticalBurst
req, ok := s.nextRequest(&pendingNormal, &pendingCritical, forceNormal)
if !ok {
s.waitAdmissions()
s.failPending(s.stoppedErr(), pendingNormal, pendingCritical)
return
}
normalWaiting := s.hasNormalRequest(&pendingNormal)
if req.priority == controlPriorityCritical {
switch {
case normalWaiting && criticalBurst >= controlBatchMaxCriticalBurst:
pushControlPendingFront(req, &pendingNormal, &pendingCritical)
var available bool
req, available = s.tryNextRequest(controlPriorityNormal, &pendingNormal, &pendingCritical)
if !available {
s.waitAdmissions()
s.failPending(s.stoppedErr(), pendingNormal, pendingCritical)
return
}
forceNormal = true
case !normalWaiting:
criticalBurst = 0
}
}
batch := []controlBatchRequest{req}
drain:
for len(batch) < controlBatchMaxPayloads {
select {
case <-s.stopCh:
s.failPending(s.stoppedErr())
return
case next := <-s.reqCh:
batch = append(batch, next)
default:
break drain
batchBytes := len(req.payload)
softLimit := s.batchSoftPayloadLimit()
batchLimit := controlBatchMaxPayloads
if req.priority == controlPriorityCritical && normalWaiting {
remaining := controlBatchMaxCriticalBurst - criticalBurst
if remaining < batchLimit {
batchLimit = remaining
}
}
payloads := make([][]byte, 0, len(batch))
for _, item := range batch {
payloads = append(payloads, item.payload)
for len(batch) < batchLimit {
next, available := s.tryNextRequest(req.priority, &pendingNormal, &pendingCritical)
if !available {
break
}
if !controlBatchCanAppend(batchBytes, len(next.payload), softLimit) {
pushControlPendingFront(next, &pendingNormal, &pendingCritical)
break
}
batch = append(batch, next)
batchBytes += len(next.payload)
}
err := s.flush(payloads, controlBatchRequestsEarliestDeadline(batch))
s.flushMu.Lock()
if req.priority == controlPriorityNormal && !forceNormal && criticalBurst < controlBatchMaxCriticalBurst {
if critical, available := s.tryNextRequest(controlPriorityCritical, &pendingNormal, &pendingCritical); available {
pendingNormal = append(append(make([]controlBatchRequest, 0, len(batch)+len(pendingNormal)), batch...), pendingNormal...)
batch = []controlBatchRequest{critical}
batchBytes = len(critical.payload)
batchLimit = controlBatchMaxCriticalBurst - criticalBurst
if batchLimit > controlBatchMaxPayloads {
batchLimit = controlBatchMaxPayloads
}
for len(batch) < batchLimit {
next, ok := s.tryNextRequest(controlPriorityCritical, &pendingNormal, &pendingCritical)
if !ok {
break
}
if !controlBatchCanAppend(batchBytes, len(next.payload), softLimit) {
pushControlPendingFront(next, &pendingNormal, &pendingCritical)
break
}
batch = append(batch, next)
batchBytes += len(next.payload)
}
}
}
batchPriority := batch[0].priority
err := s.flushQueued(batch)
s.flushMu.Unlock()
if err != nil {
s.setErr(err)
for _, item := range batch {
item.done <- err
}
s.failPending(err)
s.markFailed(err)
s.waitAdmissions()
s.failPending(err, pendingNormal, pendingCritical)
return
}
for _, item := range batch {
item.done <- nil
if batchPriority == controlPriorityCritical {
criticalBurst += len(batch)
if criticalBurst > controlBatchMaxCriticalBurst {
criticalBurst = controlBatchMaxCriticalBurst
}
} else {
criticalBurst = 0
}
}
}
func (s *controlBatchSender) nextRequest() (controlBatchRequest, bool) {
select {
case <-s.stopCh:
s.failPending(s.stoppedErr())
return controlBatchRequest{}, false
case req := <-s.reqCh:
func (s *controlBatchSender) nextRequest(pendingNormal *[]controlBatchRequest, pendingCritical *[]controlBatchRequest, forceNormal bool) (controlBatchRequest, bool) {
if forceNormal {
if req, ok := popControlPending(pendingNormal); ok {
return req, true
}
select {
case req := <-s.normalCh:
return req, true
default:
}
}
if req, ok := popControlPending(pendingCritical); ok {
return req, true
}
select {
case req := <-s.criticalCh:
return req, true
default:
}
if req, ok := popControlPending(pendingNormal); ok {
return req, true
}
select {
case <-s.stopCh:
return controlBatchRequest{}, false
case req := <-s.criticalCh:
return req, true
case req := <-s.normalCh:
return req, true
}
}
func (s *controlBatchSender) tryNextRequest(priority controlPriority, pendingNormal *[]controlBatchRequest, pendingCritical *[]controlBatchRequest) (controlBatchRequest, bool) {
pending := pendingNormal
if priority == controlPriorityCritical {
pending = pendingCritical
}
if req, ok := popControlPending(pending); ok {
return req, true
}
select {
case <-s.stopCh:
return controlBatchRequest{}, false
case req := <-s.requestChannel(priority):
return req, true
default:
return controlBatchRequest{}, false
}
}
func popControlPending(pending *[]controlBatchRequest) (controlBatchRequest, bool) {
if pending == nil || len(*pending) == 0 {
return controlBatchRequest{}, false
}
req := (*pending)[0]
*pending = (*pending)[1:]
return req, true
}
func pushControlPendingFront(req controlBatchRequest, pendingNormal *[]controlBatchRequest, pendingCritical *[]controlBatchRequest) {
pending := pendingNormal
if req.priority == controlPriorityCritical {
pending = pendingCritical
}
*pending = append([]controlBatchRequest{req}, (*pending)...)
}
func (s *controlBatchSender) requestChannel(priority controlPriority) chan controlBatchRequest {
if priority == controlPriorityCritical {
return s.criticalCh
}
return s.normalCh
}
func (s *controlBatchSender) hasNormalRequest(pendingNormal *[]controlBatchRequest) bool {
if pendingNormal != nil && len(*pendingNormal) > 0 {
return true
}
return s != nil && len(s.normalCh) > 0
}
func controlBatchCanAppend(batchBytes int, nextBytes int, softLimit int) bool {
if batchBytes == 0 {
return true
}
if softLimit <= 0 {
return false
}
return nextBytes <= softLimit-batchBytes
}
func (s *controlBatchSender) batchSoftPayloadLimit() int {
if s == nil || s.binding == nil {
return controlAdaptiveSoftPayloadFallbackBytes
}
return s.binding.controlAdaptiveSoftPayloadBytesSnapshot()
}
func (s *controlBatchSender) reserveQueueBytes(ctx context.Context, size int64, priority controlPriority) error {
limit := int64(controlBatchMaxQueuedBytes)
if priority == controlPriorityCritical {
limit += int64(controlBatchCriticalReservedBytes)
}
for {
current := s.queuedSize.Load()
if (current == 0 && size > limit) || size <= limit-current {
if s.queuedSize.CompareAndSwap(current, current+size) {
return nil
}
continue
}
select {
case <-ctx.Done():
return newTransportSendError(TransportSendStageQueue, ctx.Err())
case <-s.stopCh:
return s.stoppedErr()
case <-s.budgetWake:
}
}
}
func (s *controlBatchSender) releaseQueueBytes(size int64) {
if s == nil || size <= 0 {
return
}
s.queuedSize.Add(-size)
select {
case s.budgetWake <- struct{}{}:
default:
}
}
func (s *controlBatchSender) flushDirect(req controlBatchRequest) error {
if s == nil || s.binding == nil {
return newTransportSendError(TransportSendStageTransport, errTransportDetached)
}
queue := s.binding.queueSnapshot()
if queue == nil {
return newTransportSendError(TransportSendStageTransport, errTransportFrameQueueUnavailable)
}
var preWriteErr error
didWrite := false
started := time.Now()
lockAcquired, err := s.binding.withConnWriteLockContextStopDeadlineManaged(req.ctx, s.stopCh, writeDeadlineFromTimeout(req.writeTime), func(conn net.Conn) error {
if stoppedErr := s.errSnapshot(); stoppedErr != nil {
preWriteErr = stoppedErr
return nil
}
if ctxErr := req.contextErr(); ctxErr != nil {
preWriteErr = ctxErr
return nil
}
didWrite = true
return writeFramedPayloadBatchUnlocked(conn, queue, [][]byte{req.payload})
})
if preWriteErr != nil {
return preWriteErr
}
if !lockAcquired {
if ctxErr := req.contextErr(); ctxErr != nil {
return ctxErr
}
if stoppedErr := s.errSnapshot(); stoppedErr != nil {
return stoppedErr
}
return newTransportSendError(TransportSendStageTransport, err)
}
if didWrite {
s.binding.observeControlAdaptivePayloadWrite(len(req.payload), time.Since(started), req.writeTime, err)
}
if err != nil {
if lockAcquired {
// A framed write may have left a partial frame on the wire. Do not
// let a later request reuse this connection.
s.binding.closeConn()
}
return newTransportSendError(TransportSendStageWrite, err)
}
return nil
}
func (s *controlBatchSender) flushQueued(requests []controlBatchRequest) error {
if s == nil || s.binding == nil {
err := newTransportSendError(TransportSendStageTransport, errTransportDetached)
s.finishRequests(requests, err)
return err
}
queue := s.binding.queueSnapshot()
if queue == nil {
err := newTransportSendError(TransportSendStageTransport, errTransportFrameQueueUnavailable)
s.finishRequests(requests, err)
return err
}
pending := requests
for len(pending) > 0 {
pending = s.finishCanceledRequests(pending)
if len(pending) == 0 {
return nil
}
if stoppedErr := s.errSnapshot(); stoppedErr != nil {
s.finishRequests(pending, stoppedErr)
return stoppedErr
}
waitCtx, cleanupWait := s.controlBatchWaitContext(pending)
active := make([]controlBatchRequest, 0, len(pending))
var preWriteErr error
payloadBytes := 0
didWrite := false
started := time.Now()
writeDeadline := earlierWriteDeadline(
writeDeadlineFromTimeout(controlBatchRequestsShortestWriteTimeout(pending)),
controlBatchRequestsEarliestDeadline(pending),
)
lockAcquired, err := s.binding.withConnWriteLockContextStopDeadlineManaged(waitCtx, s.stopCh, writeDeadline, func(conn net.Conn) error {
if stoppedErr := s.errSnapshot(); stoppedErr != nil {
preWriteErr = stoppedErr
return nil
}
for _, item := range pending {
if cancelErr, canceled := item.cancelBeforeStart(); canceled {
s.finishRequest(item, cancelErr)
continue
}
if !item.tryStart() {
s.finishRequest(item, item.canceledErr())
continue
}
active = append(active, item)
payloadBytes += len(item.payload)
}
if len(active) == 0 {
return nil
}
payloads := make([][]byte, 0, len(active))
for _, item := range active {
payloads = append(payloads, item.payload)
}
didWrite = true
return writeFramedPayloadBatchUnlocked(conn, queue, payloads)
})
cleanupWait()
if preWriteErr != nil {
s.finishRequests(pending, preWriteErr)
return preWriteErr
}
if !lockAcquired {
remaining := s.finishCanceledRequests(pending)
if len(remaining) < len(pending) {
pending = remaining
continue
}
if stoppedErr := s.errSnapshot(); stoppedErr != nil {
s.finishRequests(pending, stoppedErr)
return stoppedErr
}
transportErr := newTransportSendError(TransportSendStageTransport, err)
s.finishRequests(pending, transportErr)
return transportErr
}
if len(active) == 0 {
if err == nil {
return nil
}
writeErr := newTransportSendError(TransportSendStageWrite, err)
s.finishRequests(pending, writeErr)
return writeErr
}
writeTimeout := controlBatchRequestsShortestWriteTimeout(active)
if didWrite {
s.binding.observeControlAdaptivePayloadWrite(payloadBytes, time.Since(started), writeTimeout, err)
}
if err != nil {
if lockAcquired {
// A framed write may have left a partial frame on the wire. Do not
// let a later request reuse this connection.
s.binding.closeConn()
}
err = newTransportSendError(TransportSendStageWrite, err)
}
s.finishRequests(active, err)
return err
}
return nil
}
func (s *controlBatchSender) controlBatchWaitContext(requests []controlBatchRequest) (context.Context, func()) {
base := context.Background()
if s != nil && s.stopCtx != nil {
base = s.stopCtx
}
ctx, cancel := context.WithCancel(base)
stops := make([]func() bool, 0, len(requests))
for _, item := range requests {
if item.ctx == nil || item.ctx.Done() == nil {
continue
}
stops = append(stops, context.AfterFunc(item.ctx, cancel))
}
return ctx, func() {
for _, stop := range stops {
stop()
}
cancel()
}
}
func controlBatchRequestsShortestWriteTimeout(batch []controlBatchRequest) time.Duration {
var timeout time.Duration
for _, item := range batch {
if item.writeTime <= 0 {
continue
}
if timeout == 0 || item.writeTime < timeout {
timeout = item.writeTime
}
}
return timeout
}
func controlBatchRequestsEarliestDeadline(batch []controlBatchRequest) time.Time {
var deadline time.Time
for _, item := range batch {
if item.deadline.IsZero() {
continue
}
if deadline.IsZero() || item.deadline.Before(deadline) {
deadline = item.deadline
}
candidate := contextDeadline(item.ctx)
deadline = earlierWriteDeadline(deadline, candidate)
}
return deadline
}
func (s *controlBatchSender) flush(payloads [][]byte, deadline time.Time) error {
if s == nil || s.binding == nil {
return errTransportDetached
func (s *controlBatchSender) finishCanceledRequests(requests []controlBatchRequest) []controlBatchRequest {
remaining := requests[:0]
for _, item := range requests {
if cancelErr, canceled := item.cancelBeforeStart(); canceled {
s.finishRequest(item, cancelErr)
continue
}
remaining = append(remaining, item)
}
queue := s.binding.queueSnapshot()
if queue == nil {
return errTransportFrameQueueUnavailable
return remaining
}
func (s *controlBatchSender) finishRequests(requests []controlBatchRequest, err error) {
for _, item := range requests {
s.finishRequest(item, err)
}
}
func (s *controlBatchSender) finishRequest(req controlBatchRequest, err error) {
if s != nil {
s.queued.Add(-1)
s.releaseQueueBytes(req.queueSize)
}
req.done <- err
}
func (s *controlBatchSender) beginAdmission() bool {
if s == nil {
return false
}
s.admissionMu.Lock()
defer s.admissionMu.Unlock()
if s.admissionClosed {
return false
}
s.admitting.Add(1)
return true
}
func (s *controlBatchSender) endAdmission() {
if s != nil {
s.admitting.Done()
}
}
func (s *controlBatchSender) waitAdmissions() {
if s != nil {
s.admitting.Wait()
}
return s.binding.withConnWriteLockDeadline(deadline, func(conn net.Conn) error {
return writeFramedPayloadBatchUnlocked(conn, queue, payloads)
})
}
func (s *controlBatchSender) stop() {
if s == nil {
return
}
s.stopOnce.Do(func() {
s.setErr(errTransportDetached)
close(s.stopCh)
})
s.markFailed(newTransportSendError(TransportSendStageTransport, errTransportDetached))
<-s.doneCh
s.flushMu.Lock()
s.flushMu.Unlock()
}
func (s *controlBatchSender) failPending(err error) {
func (s *controlBatchSender) failPending(err error, pendingNormal []controlBatchRequest, pendingCritical []controlBatchRequest) {
for _, item := range pendingCritical {
s.finishRequest(item, err)
}
for _, item := range pendingNormal {
s.finishRequest(item, err)
}
for {
select {
case item := <-s.reqCh:
item.done <- err
case item := <-s.criticalCh:
s.finishRequest(item, err)
case item := <-s.normalCh:
s.finishRequest(item, err)
default:
return
}
@@ -164,9 +739,25 @@ func (s *controlBatchSender) setErr(err error) {
s.errMu.Unlock()
}
func (s *controlBatchSender) markFailed(err error) {
if s == nil {
return
}
s.setErr(err)
s.stopOnce.Do(func() {
s.admissionMu.Lock()
s.admissionClosed = true
if s.stopCancel != nil {
s.stopCancel()
}
close(s.stopCh)
s.admissionMu.Unlock()
})
}
func (s *controlBatchSender) errSnapshot() error {
if s == nil {
return errTransportDetached
return newTransportSendError(TransportSendStageTransport, errTransportDetached)
}
s.errMu.Lock()
defer s.errMu.Unlock()
@@ -177,5 +768,69 @@ func (s *controlBatchSender) stoppedErr() error {
if err := s.errSnapshot(); err != nil {
return err
}
return errTransportDetached
return newTransportSendError(TransportSendStageTransport, errTransportDetached)
}
func (r controlBatchRequest) tryStart() bool {
if r.state == nil {
return true
}
return r.state.value.CompareAndSwap(controlBatchRequestQueued, controlBatchRequestStarted)
}
func (r controlBatchRequest) tryCancel() bool {
return r.state != nil && r.state.value.CompareAndSwap(controlBatchRequestQueued, controlBatchRequestCanceled)
}
func (r controlBatchRequest) cancelBeforeStart() (error, bool) {
if r.state == nil {
if err := r.contextErr(); err != nil {
return err, true
}
return nil, false
}
if r.state.value.Load() == controlBatchRequestCanceled {
return r.canceledErr(), true
}
if err := r.contextErr(); err != nil && r.tryCancel() {
return err, true
}
if r.state.value.Load() == controlBatchRequestCanceled {
return r.canceledErr(), true
}
return nil, false
}
func (r controlBatchRequest) contextErr() error {
if r.ctx == nil {
return nil
}
err := r.ctx.Err()
if err == nil {
if deadline, ok := r.ctx.Deadline(); ok && !time.Now().Before(deadline) {
err = context.DeadlineExceeded
}
}
return newTransportSendError(TransportSendStageQueue, err)
}
func (r controlBatchRequest) canceledErr() error {
if r.ctx != nil && r.ctx.Err() != nil {
return newTransportSendError(TransportSendStageQueue, r.ctx.Err())
}
return newTransportSendError(TransportSendStageQueue, context.Canceled)
}
func maxInt(left int, right int) int {
if left > right {
return left
}
return right
}
func maxDuration(left time.Duration, right time.Duration) time.Duration {
if left > right {
return left
}
return right
}
+33 -5
View File
@@ -2,13 +2,16 @@ package notify
import (
"b612.me/notify/internal/timeutil"
"context"
crand "crypto/rand"
"encoding/binary"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"sync/atomic"
"time"
)
type EnvelopeKind uint8
@@ -30,6 +33,11 @@ type Envelope struct {
Body []byte
Stream StreamPacket
File FilePacket
controlCtx context.Context
controlPriority controlPriority
controlTimeout time.Duration
transportProfile *transportProtectionProfile
}
type StreamPacket struct {
@@ -56,12 +64,31 @@ func wrapTransferMsgEnvelope(msg TransferMsg, enFn func(interface{}) ([]byte, er
return Envelope{}, err
}
return Envelope{
Kind: EnvelopeSignal,
ID: msg.ID,
Body: body,
Kind: EnvelopeSignal,
ID: msg.ID,
Body: body,
controlPriority: controlPriorityForTransferMessage(msg),
}, nil
}
func controlPriorityForTransferMessage(msg TransferMsg) controlPriority {
switch msg.Type {
case MSG_SYS, MSG_SYS_WAIT, MSG_SYS_REPLY, MSG_KEY_CHANGE, MSG_SYNC_REPLY:
return controlPriorityCritical
}
if msg.Key == "heartbeat" || msg.Key == "bye" || strings.HasPrefix(msg.Key, "notify.") {
return controlPriorityCritical
}
return controlPriorityNormal
}
func (env Envelope) controlContext() context.Context {
if env.controlCtx == nil {
return context.Background()
}
return env.controlCtx
}
func unwrapTransferMsgEnvelope(env Envelope, deFn func([]byte) (interface{}, error)) (TransferMsg, error) {
if env.Kind != EnvelopeSignal {
return TransferMsg{}, errors.New("envelope kind is not signal")
@@ -79,8 +106,9 @@ func unwrapTransferMsgEnvelope(env Envelope, deFn func([]byte) (interface{}, err
func newSignalAckEnvelope(signalID uint64) Envelope {
return Envelope{
Kind: EnvelopeSignalAck,
ID: signalID,
Kind: EnvelopeSignalAck,
ID: signalID,
controlPriority: controlPriorityCritical,
}
}
+17 -12
View File
@@ -9,14 +9,15 @@ import (
)
type LogicalConn struct {
client *ClientConn
server Server
ClientID string
ClientAddr net.Addr
state atomic.Pointer[logicalConnState]
runtime atomic.Pointer[logicalConnRuntimeState]
transportState atomic.Pointer[clientConnTransportState]
attachment atomic.Pointer[clientConnAttachmentState]
client *ClientConn
server Server
ClientID string
ClientAddr net.Addr
state atomic.Pointer[logicalConnState]
runtime atomic.Pointer[logicalConnRuntimeState]
transportState atomic.Pointer[clientConnTransportState]
attachment atomic.Pointer[clientConnAttachmentState]
inboundTransitionProfile atomic.Pointer[transportProtectionProfile]
}
var errLogicalConnClientNil = errors.New("logical conn is nil")
@@ -856,6 +857,10 @@ func (c *LogicalConn) sessionRuntimeSnapshot() *clientConnSessionRuntime {
}
func (c *LogicalConn) setSessionRuntime(rt *clientConnSessionRuntime) {
c.setSessionRuntimeWithCloseOld(rt, false)
}
func (c *LogicalConn) setSessionRuntimeWithCloseOld(rt *clientConnSessionRuntime, closeOld bool) {
if c == nil || rt == nil {
return
}
@@ -883,7 +888,7 @@ func (c *LogicalConn) setSessionRuntime(rt *clientConnSessionRuntime) {
client.syncLegacySessionRuntimeFromState(state)
}
if oldBinding != nil {
oldBinding.stopBackgroundWorkers()
stopReplacedTransportBinding(oldBinding, rt.transport, closeOld)
}
}
@@ -982,14 +987,14 @@ func (c *LogicalConn) startSession(tuConn net.Conn, stopCtx context.Context, sto
transportGeneration = c.markTransportAttached()
c.clearTransportDetachState()
}
c.setSessionRuntime(&clientConnSessionRuntime{
c.setSessionRuntimeWithCloseOld(&clientConnSessionRuntime{
transport: newTransportBinding(tuConn, nil),
transportAttached: tuConn != nil,
transportGeneration: transportGeneration,
tuConn: tuConn,
stopCtx: stopCtx,
stopFn: stopFn,
})
}, true)
c.markSessionStarted()
return stopCtx, stopFn
}
@@ -1030,7 +1035,7 @@ func (c *LogicalConn) attachSessionTransport(tuConn net.Conn) error {
next.transportStopCtx = nil
next.transportStopFn = nil
next.transportDone = nil
c.setSessionRuntime(&next)
c.setSessionRuntimeWithCloseOld(&next, true)
if tuConn.RemoteAddr() != nil {
c.setRemoteAddr(tuConn.RemoteAddr())
}
+61 -20
View File
@@ -1,10 +1,13 @@
package notify
import (
"context"
"net"
"time"
)
const defaultMessageReplyWriteTimeout = 30 * time.Second
const (
MSG_SYS MessageType = iota
MSG_SYS_WAIT
@@ -55,16 +58,40 @@ type WaitMsg struct {
}
type messageLogicalTransferSender interface {
sendLogical(*LogicalConn, TransferMsg) (WaitMsg, error)
sendLogicalContext(context.Context, *LogicalConn, TransferMsg, time.Duration) (WaitMsg, error)
}
type messageTransportTransferSender interface {
sendTransportContextWithWriteTimeout(context.Context, *TransportConn, TransferMsg, time.Duration) (WaitMsg, error)
}
type messageInboundTransferSender interface {
sendTransferInbound(*LogicalConn, *TransportConn, net.Conn, *transportProtectionProfile, TransferMsg) error
sendTransferInboundContext(context.Context, *LogicalConn, *TransportConn, net.Conn, *transportProtectionProfile, TransferMsg, time.Duration) error
}
type messageClientTransferSender interface {
sendWithContextTimeout(context.Context, TransferMsg, time.Duration) (WaitMsg, error)
}
type messageReplyWriteTimeoutProvider interface {
ReplyWriteTimeout() time.Duration
}
func (m *Message) Reply(value MsgVal) (err error) {
return m.replyContext(context.Background(), value)
}
func (m *Message) ReplyCtx(ctx context.Context, value MsgVal) (err error) {
if ctx == nil {
ctx = context.Background()
}
return m.replyContext(ctx, value)
}
func (m *Message) replyContext(ctx context.Context, value MsgVal) (err error) {
logical := messageLogicalConnSnapshot(m)
transport := messageTransportConnSnapshot(m)
writeTimeout := defaultMessageReplyWriteTimeout
reply := TransferMsg{
ID: m.ID,
Key: m.Key,
@@ -78,21 +105,6 @@ func (m *Message) Reply(value MsgVal) (err error) {
reply.Type = MSG_SYS_REPLY
}
if m.NetType == NET_SERVER {
if m.inboundConn != nil && logical != nil {
server := logical.Server()
if server == nil {
return transportDetachedErrorForPeer(logical, transport)
}
sender, _ := server.(messageInboundTransferSender)
if sender == nil {
return transportDetachedErrorForPeer(logical, transport)
}
return sender.sendTransferInbound(logical, transport, m.inboundConn, messageInboundTransportProtectionSnapshot(m), reply)
}
if transport != nil {
_, err = transport.sendTransfer(reply)
return
}
if logical == nil {
return transportDetachedErrorForPeer(nil, transport)
}
@@ -100,24 +112,53 @@ func (m *Message) Reply(value MsgVal) (err error) {
if server == nil {
return transportDetachedErrorForPeer(logical, transport)
}
if provider, ok := server.(messageReplyWriteTimeoutProvider); ok {
writeTimeout = provider.ReplyWriteTimeout()
}
if m.inboundConn != nil && logical != nil {
sender, _ := server.(messageInboundTransferSender)
if sender == nil {
return transportDetachedErrorForPeer(logical, transport)
}
return sender.sendTransferInboundContext(ctx, logical, transport, m.inboundConn, messageInboundTransportProtectionSnapshot(m), reply, writeTimeout)
}
if transport != nil {
sender, _ := server.(messageTransportTransferSender)
if sender == nil {
return transportDetachedErrorForPeer(logical, transport)
}
_, err = sender.sendTransportContextWithWriteTimeout(ctx, transport, reply, writeTimeout)
return err
}
sender, _ := server.(messageLogicalTransferSender)
if sender == nil {
return transportDetachedErrorForPeer(logical, transport)
}
_, err = sender.sendLogical(logical, reply)
_, err = sender.sendLogicalContext(ctx, logical, reply, writeTimeout)
}
if m.NetType == NET_CLIENT {
_, err = m.ServerConn.send(reply)
if m.ServerConn == nil {
return net.ErrClosed
}
if sender, ok := m.ServerConn.(messageClientTransferSender); ok {
_, err = sender.sendWithContextTimeout(ctx, reply, writeTimeout)
} else {
_, err = m.ServerConn.send(reply)
}
}
return
}
func (m *Message) ReplyObj(value interface{}) (err error) {
return m.ReplyObjCtx(context.Background(), value)
}
func (m *Message) ReplyObjCtx(ctx context.Context, value interface{}) (err error) {
data, err := encode(value)
if err != nil {
return err
}
return m.Reply(data)
return m.ReplyCtx(ctx, data)
}
func hydrateServerMessagePeerFields(message Message) Message {
+6 -3
View File
@@ -20,10 +20,12 @@ func newRunningPeerAttachServerForTest(t *testing.T, configure func(*ServerCommo
}
stopCtx, stopFn := context.WithCancel(context.Background())
queue := stario.NewQueueCtx(stopCtx, 8, math.MaxUint32)
inboundDispatcher := newInboundDispatcher()
server.setServerSessionRuntime(&serverSessionRuntime{
stopCtx: stopCtx,
stopFn: stopFn,
queue: queue,
stopCtx: stopCtx,
stopFn: stopFn,
queue: queue,
inboundDispatcher: inboundDispatcher,
})
server.markSessionStarted()
@@ -33,6 +35,7 @@ func newRunningPeerAttachServerForTest(t *testing.T, configure func(*ServerCommo
t.Cleanup(func() {
transportStop()
stopFn()
inboundDispatcher.CloseAndWait()
})
return server
}
+7
View File
@@ -63,6 +63,13 @@ func transportDetachedError(detail string, cause error) error {
return newDetailedStateError(errTransportDetached, detail, cause)
}
// IsTransportDetachedError reports whether err means that the physical
// transport was detached. The result remains true through the detailed and
// transport-send wrappers used by the public send/read APIs.
func IsTransportDetachedError(err error) bool {
return err != nil && errors.Is(err, errTransportDetached)
}
func clientTransportDetachedError(c *ClientCommon) error {
if c == nil {
return errTransportDetached
+37
View File
@@ -0,0 +1,37 @@
package notify
import (
"errors"
"fmt"
"io"
"testing"
)
func TestIsTransportDetachedErrorRecognizesWrappedDetach(t *testing.T) {
base := transportDetachedError("dedicated bulk read error", io.ErrUnexpectedEOF)
for name, err := range map[string]error{
"direct": base,
"wrapped": fmt.Errorf("bulk reset: %w", base),
"transport-send": newTransportSendError(TransportSendStageTransport, base),
} {
t.Run(name, func(t *testing.T) {
if !IsTransportDetachedError(err) {
t.Fatalf("IsTransportDetachedError(%v) = false", err)
}
})
}
}
func TestIsTransportDetachedErrorRejectsOtherErrors(t *testing.T) {
for name, err := range map[string]error{
"nil": nil,
"business": errors.New("remote file became a directory"),
"eof": io.EOF,
} {
t.Run(name, func(t *testing.T) {
if IsTransportDetachedError(err) {
t.Fatalf("IsTransportDetachedError(%v) = true", err)
}
})
}
}
+28 -9
View File
@@ -10,8 +10,9 @@ import (
)
const (
systemPeerAttachKey = "_notify_peer_attach"
peerAttachTimeout = 5 * time.Second
systemPeerAttachKey = "_notify_peer_attach"
peerAttachTimeout = 5 * time.Second
peerAttachTransitionFallbackTTL = peerAttachTimeout
)
type peerAttachRequest struct {
@@ -206,11 +207,17 @@ func (s *ServerCommon) replyPeerAttach(client *LogicalConn, message Message, res
Value: encoded,
Type: MSG_SYS_REPLY,
}
transport := messageTransportConnSnapshot(&message)
profile := messageInboundTransportProtectionSnapshot(&message)
if message.inboundConn != nil {
return s.sendTransferInbound(client, messageTransportConnSnapshot(&message), message.inboundConn, messageInboundTransportProtectionSnapshot(&message), reply)
return s.sendTransferInbound(client, transport, message.inboundConn, profile, reply)
}
_, err = s.sendLogical(client, reply)
return err
env, err := wrapTransferMsgEnvelope(reply, s.sequenceEn)
if err != nil {
return err
}
env.transportProfile = profile
return s.sendSignalEnvelopeMaybeReliableTransport(transport, env, reply)
}
func (s *ServerCommon) handlePeerAttachSystemMessage(message Message) bool {
@@ -219,6 +226,10 @@ func (s *ServerCommon) handlePeerAttachSystemMessage(message Message) bool {
}
message = hydrateServerMessagePeerFields(message)
current := messageLogicalConnSnapshot(&message)
if message.inboundTransportProfile == nil && current != nil {
profile := current.transportProtectionProfileSnapshot()
message.inboundTransportProfile = &profile
}
transport := message.inboundConn
if transport == nil && current != nil {
transport = current.transportSnapshot()
@@ -281,12 +292,20 @@ func (s *ServerCommon) handlePeerAttachSystemMessage(message Message) bool {
s.peerAttachAuthFallbackCount.Add(1)
}
}
if err := s.replyPeerAttach(bound, message, resp); err != nil && bound != nil {
s.stopLogicalSession(bound, "peer attach reply failed", err)
return true
}
var transitionProfile *transportProtectionProfile
if bound != nil && s.securityConfigured {
if message.inboundTransportProfile != nil {
transitionProfile = bound.installInboundTransitionProfile(*message.inboundTransportProfile)
}
bound.applyTransportProtectionProfile(steadyProfile)
}
replyErr := s.replyPeerAttach(bound, message, resp)
if transitionProfile != nil {
bound.clearInboundTransitionProfile(transitionProfile)
}
if replyErr != nil && bound != nil {
s.stopLogicalSession(bound, "peer attach reply failed", replyErr)
return true
}
return true
}
+129 -1
View File
@@ -1,6 +1,7 @@
package notify
import (
"b612.me/stario"
"net"
"testing"
"time"
@@ -10,6 +11,13 @@ func TestClientPeerAttachRenamesAcceptedPeer(t *testing.T) {
secret := []byte("0123456789abcdef0123456789abcdef")
server := newRunningPeerAttachServerForTest(t, func(server *ServerCommon) {
server.SetSecretKey(secret)
if err := UseSignalReliabilityServer(server, &SignalReliabilityOptions{
Enabled: true,
AckTimeout: 200 * time.Millisecond,
SendRetry: 2,
}); err != nil {
t.Fatal(err)
}
})
client := NewClient().(*ClientCommon)
client.SetSecretKey(secret)
@@ -113,7 +121,6 @@ func TestReplyPeerAttachUsesInboundConnWithoutWaitingSignalAck(t *testing.T) {
t.Fatalf("UseSignalReliabilityServer failed: %v", err)
}
})
clientConn, serverConn := net.Pipe()
defer clientConn.Close()
defer serverConn.Close()
@@ -176,3 +183,124 @@ func TestReplyPeerAttachUsesInboundConnWithoutWaitingSignalAck(t *testing.T) {
t.Fatalf("reply key = %q, want %q", transfer.Key, systemPeerAttachKey)
}
}
func TestReplyPeerAttachUsesCapturedProfileWithoutInboundConn(t *testing.T) {
secret := []byte("0123456789abcdef0123456789abcdef")
server := newRunningPeerAttachServerForTest(t, func(server *ServerCommon) {
server.SetSecretKey(secret)
if err := UseSignalReliabilityServer(server, &SignalReliabilityOptions{
Enabled: true,
AckTimeout: 2 * time.Second,
SendRetry: 2,
}); err != nil {
t.Fatal(err)
}
})
clientConn, serverConn := net.Pipe()
defer clientConn.Close()
defer serverConn.Close()
logical := bootstrapPeerAttachLogicalForTest(t, server, serverConn)
originalProfile := logical.transportProtectionProfileSnapshot()
message := Message{
NetType: NET_SERVER,
LogicalConn: logical,
TransportConn: logical.CurrentTransportConn(),
TransferMsg: TransferMsg{
ID: 43,
Key: systemPeerAttachKey,
Type: MSG_SYS_WAIT,
},
Time: time.Now(),
inboundTransportProfile: &originalProfile,
}
alternate, err := deriveModernPSKProtectionProfile([]byte("notify-peer-attach-no-inbound-conn"), testModernPSKOptions(), ProtectionManaged)
if err != nil {
t.Fatalf("deriveModernPSKProtectionProfile failed: %v", err)
}
logical.applyTransportProtectionProfile(alternate)
transition := logical.installInboundTransitionProfile(originalProfile)
defer logical.clearInboundTransitionProfile(transition)
done := make(chan error, 1)
go func() {
done <- server.replyPeerAttach(logical, message, peerAttachResponse{
PeerID: "peer-test",
Accepted: true,
})
}()
env := readServerEnvelopeFromConnWithProfile(t, server, originalProfile, clientConn, time.Second)
if env.Kind != EnvelopeSignal {
t.Fatalf("reply envelope kind = %v, want %v", env.Kind, EnvelopeSignal)
}
ackPlain, err := server.encodeEnvelopePlain(newSignalAckEnvelope(env.ID))
if err != nil {
t.Fatal(err)
}
ackPayload, err := encryptTransportPayloadCodec(originalProfile.mode, originalProfile.runtime, originalProfile.msgEn, originalProfile.secretKey, ackPlain)
if err != nil {
t.Fatal(err)
}
if err := writeFullToConn(clientConn, stario.NewQueue().BuildMessage(ackPayload)); err != nil {
t.Fatalf("write bootstrap signal ack: %v", err)
}
select {
case err := <-done:
if err != nil {
t.Fatalf("replyPeerAttach failed: %v", err)
}
case <-time.After(time.Second):
t.Fatal("replyPeerAttach should finish without an inbound stream conn")
}
}
func TestPeerAttachTransitionProfileDecryptsBootstrapFramesAndClears(t *testing.T) {
bootstrap, err := deriveModernPSKProtectionProfile([]byte("notify-peer-transition-bootstrap"), testModernPSKOptions(), ProtectionManaged)
if err != nil {
t.Fatal(err)
}
steady, err := deriveModernPSKProtectionProfile([]byte("notify-peer-transition-steady"), testModernPSKOptions(), ProtectionManaged)
if err != nil {
t.Fatal(err)
}
payload, err := encryptTransportPayloadCodec(bootstrap.mode, bootstrap.runtime, bootstrap.msgEn, bootstrap.secretKey, []byte("bootstrap-frame"))
if err != nil {
t.Fatal(err)
}
t.Run("client", func(t *testing.T) {
client := NewClient().(*ClientCommon)
client.setClientTransportProtectionProfile(steady)
transition := client.installInboundTransitionProfile(bootstrap)
plain, release, err := client.decryptTransportPayloadPooled(append([]byte(nil), payload...), nil)
if release != nil {
release()
}
if err != nil || string(plain) != "bootstrap-frame" {
t.Fatalf("client transition decrypt plain=%q err=%v", plain, err)
}
client.clearInboundTransitionProfile(transition)
if _, _, err := client.decryptTransportPayloadPooled(append([]byte(nil), payload...), nil); err == nil {
t.Fatal("client accepted bootstrap frame after transition fallback cleared")
}
})
t.Run("server", func(t *testing.T) {
logical := newServerLogicalConn(nil, "transition-server", nil)
logical.applyTransportProtectionProfile(steady)
transition := logical.installInboundTransitionProfile(bootstrap)
plain, release, err := (&ServerCommon{}).decryptTransportPayloadLogicalPooled(logical, append([]byte(nil), payload...), nil)
if release != nil {
release()
}
if err != nil || string(plain) != "bootstrap-frame" {
t.Fatalf("server transition decrypt plain=%q err=%v", plain, err)
}
logical.clearInboundTransitionProfile(transition)
if _, _, err := (&ServerCommon{}).decryptTransportPayloadLogicalPooled(logical, append([]byte(nil), payload...), nil); err == nil {
t.Fatal("server accepted bootstrap frame after transition fallback cleared")
}
})
}
+279
View File
@@ -0,0 +1,279 @@
package notify
import (
"context"
"errors"
"math"
"net"
"os"
"testing"
"time"
"b612.me/stario"
)
func newServerBlackholeTransport(t *testing.T, id string) (*ServerCommon, *LogicalConn, *TransportConn, net.Conn, net.Conn) {
t.Helper()
server := NewServer().(*ServerCommon)
UseLegacySecurityServer(server)
stopCtx, stopFn := context.WithCancel(context.Background())
server.setServerSessionRuntime(&serverSessionRuntime{
stopCtx: stopCtx,
stopFn: stopFn,
queue: stario.NewQueueCtx(stopCtx, 4, math.MaxUint32),
})
server.markSessionStarted()
left, right := net.Pipe()
logical, _, _ := newRegisteredServerLogicalForTest(t, server, id, left, stopCtx, stopFn)
logical.applyAttachmentProfile(0, 0, server.defaultMsgEn, server.defaultMsgDe, server.defaultFastStreamEncode, server.defaultFastBulkEncode, server.defaultFastPlainEncode, server.handshakeRsaKey, server.SecretKey)
transport := logical.CurrentTransportConn()
if transport == nil {
t.Fatal("server transport is nil")
}
t.Cleanup(func() {
_ = left.Close()
_ = right.Close()
server.markSessionStopped("test done", nil)
})
return server, logical, transport, left, right
}
func requireBoundedServerWrite(t *testing.T, timeout time.Duration, write func(context.Context) error) {
t.Helper()
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
done := make(chan error, 1)
started := time.Now()
go func() { done <- write(ctx) }()
select {
case err := <-done:
if err == nil {
t.Fatal("blackhole server write returned nil")
}
if !errors.Is(err, context.DeadlineExceeded) && !errors.Is(err, os.ErrDeadlineExceeded) {
var netErr net.Error
if !errors.As(err, &netErr) || !netErr.Timeout() {
t.Fatalf("blackhole server write error=%v, want deadline error", err)
}
}
if elapsed := time.Since(started); elapsed > 5*timeout {
t.Fatalf("blackhole server write returned after %v, want bounded by %v", elapsed, timeout)
}
case <-time.After(10 * timeout):
t.Fatalf("blackhole server write ignored its %v context deadline", timeout)
}
}
func TestServerSharedBulkWriteHonorsContextDeadline(t *testing.T) {
server, logical, transport, _, _ := newServerBlackholeTransport(t, "shared-bulk-write-context")
requireBoundedServerWrite(t, 30*time.Millisecond, func(ctx context.Context) error {
return server.sendFastBulkDataTransport(ctx, logical, transport, 1, 0, []byte("bulk"), bulkFastPathVersionV1)
})
}
func TestServerSharedStreamWriteHonorsContextDeadline(t *testing.T) {
server, logical, transport, _, _ := newServerBlackholeTransport(t, "shared-stream-write-context")
stream := newStreamHandle(context.Background(), nil, serverFileScope(logical), StreamOpenRequest{
StreamID: "shared-stream-write-context",
DataID: 1,
Channel: StreamDataChannel,
FastPathVersion: streamFastPathVersionV1,
}, 0, logical, transport, transport.TransportGeneration(), nil, nil, nil, defaultStreamConfig())
requireBoundedServerWrite(t, 30*time.Millisecond, func(ctx context.Context) error {
return server.sendFastStreamDataTransport(ctx, logical, transport, stream, []byte("stream"))
})
}
func TestMessageReplyUsesConfiguredDefaultWriteTimeout(t *testing.T) {
server, logical, transport, left, _ := newServerBlackholeTransport(t, "reply-default-write-timeout")
server.SetReplyWriteTimeout(35 * time.Millisecond)
message := Message{
NetType: NET_SERVER,
LogicalConn: logical,
TransportConn: transport,
TransferMsg: TransferMsg{
ID: 1,
Key: "reply-default-write-timeout",
Type: MSG_SYNC_ASK,
},
inboundConn: left,
}
started := time.Now()
err := message.Reply([]byte("reply"))
if err == nil {
t.Fatal("blackhole Message.Reply returned nil")
}
if elapsed := time.Since(started); elapsed > 250*time.Millisecond {
t.Fatalf("Message.Reply returned after %v, want configured default write bound", elapsed)
}
}
func TestMessageReplyCtxUsesEarlierCallerDeadline(t *testing.T) {
server, logical, transport, left, _ := newServerBlackholeTransport(t, "reply-caller-write-timeout")
server.SetReplyWriteTimeout(time.Second)
message := Message{
NetType: NET_SERVER,
LogicalConn: logical,
TransportConn: transport,
TransferMsg: TransferMsg{
ID: 2,
Key: "reply-caller-write-timeout",
Type: MSG_SYNC_ASK,
},
inboundConn: left,
}
ctx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond)
defer cancel()
started := time.Now()
err := message.ReplyCtx(ctx, []byte("reply"))
if err == nil {
t.Fatal("blackhole Message.ReplyCtx returned nil")
}
if elapsed := time.Since(started); elapsed > 200*time.Millisecond {
t.Fatalf("Message.ReplyCtx returned after %v, want caller deadline", elapsed)
}
canceled, cancelNow := context.WithCancel(context.Background())
cancelNow()
if err := message.ReplyObjCtx(canceled, "ok"); !errors.Is(err, context.Canceled) {
t.Fatalf("Message.ReplyObjCtx canceled error=%v, want context canceled", err)
}
}
type closeTrackingWriteConn struct {
closed bool
writes int
}
func (c *closeTrackingWriteConn) Read([]byte) (int, error) { return 0, net.ErrClosed }
func (c *closeTrackingWriteConn) Close() error { c.closed = true; return nil }
func (c *closeTrackingWriteConn) LocalAddr() net.Addr { return nil }
func (c *closeTrackingWriteConn) RemoteAddr() net.Addr { return nil }
func (c *closeTrackingWriteConn) SetDeadline(time.Time) error { return nil }
func (c *closeTrackingWriteConn) SetReadDeadline(time.Time) error { return nil }
func (c *closeTrackingWriteConn) SetWriteDeadline(time.Time) error { return nil }
func (c *closeTrackingWriteConn) Write(data []byte) (int, error) {
c.writes++
return len(data), nil
}
func TestServerRawWriteGateWaitTimeoutDoesNotCloseConnection(t *testing.T) {
server, logical, transport, _, _ := newServerBlackholeTransport(t, "raw-write-gate-timeout")
conn := &closeTrackingWriteConn{}
gateRef := retainRawConnWriteGate(conn)
<-gateRef.gate
t.Cleanup(func() {
gateRef.gate <- struct{}{}
releaseRawConnWriteGate(conn, gateRef)
})
ctx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond)
defer cancel()
err := server.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, []byte("reply"), time.Second)
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("raw write gate wait error=%v, want context deadline exceeded", err)
}
if conn.closed {
t.Fatal("raw write gate wait timeout closed a connection before physical write started")
}
if conn.writes != 0 {
t.Fatalf("raw write gate wait performed %d physical writes", conn.writes)
}
}
func TestServerUDPWriteGateWaitHonorsWriteTimeout(t *testing.T) {
server := NewServer().(*ServerCommon)
stopCtx, stopFn := context.WithCancel(context.Background())
t.Cleanup(stopFn)
sender, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = sender.Close() })
receiver, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = receiver.Close() })
server.setServerSessionRuntime(&serverSessionRuntime{
stopCtx: stopCtx,
stopFn: stopFn,
queue: stario.NewQueueCtx(stopCtx, 4, math.MaxUint32),
udpListener: sender,
})
logical := newServerLogicalConn(server, "udp-write-gate-timeout", receiver.LocalAddr())
transport := &TransportConn{logical: logical, remoteAddr: receiver.LocalAddr(), attached: true}
gateHeld := make(chan struct{})
releaseGate := make(chan struct{})
firstDone := make(chan error, 1)
go func() {
firstDone <- server.withUDPWriteLock(context.Background(), func() error {
close(gateHeld)
<-releaseGate
return nil
})
}()
select {
case <-gateHeld:
case <-time.After(time.Second):
t.Fatal("first UDP writer did not acquire the write gate")
}
started := time.Now()
writeDone := make(chan error, 1)
go func() {
writeDone <- server.writeEnvelopePayloadContextTimeout(context.Background(), logical, transport, nil, []byte("reply"), 30*time.Millisecond)
}()
select {
case err = <-writeDone:
case <-time.After(200 * time.Millisecond):
close(releaseGate)
<-firstDone
<-writeDone
t.Fatal("UDP write gate wait ignored the configured write timeout")
}
if !errors.Is(err, context.DeadlineExceeded) {
close(releaseGate)
<-firstDone
t.Fatalf("UDP write gate wait error=%v, want context deadline exceeded", err)
}
if elapsed := time.Since(started); elapsed > 200*time.Millisecond {
close(releaseGate)
<-firstDone
t.Fatalf("UDP write gate wait returned after %v, want configured write bound", elapsed)
}
close(releaseGate)
if err := <-firstDone; err != nil {
t.Fatalf("first UDP writer failed: %v", err)
}
}
func TestTransportWriteGateRegistryReleasesBindingsAndRawWrites(t *testing.T) {
conn := &serializedWriteTestConn{}
first := newTransportBinding(conn, stario.NewQueue())
second := newTransportBinding(conn, stario.NewQueue())
if first.writeGateSnapshot() != second.writeGateSnapshot() {
t.Fatal("same physical connection did not share one write gate")
}
entry, ok := transportConnWriteGates.Load(conn)
if !ok {
t.Fatal("shared write gate was not registered")
}
first.stopBackgroundWorkers()
if current, ok := transportConnWriteGates.Load(conn); !ok || current != entry {
t.Fatal("stopping one shared binding removed the live write gate")
}
second.stopBackgroundWorkers()
if _, ok := transportConnWriteGates.Load(conn); ok {
t.Fatal("last binding release retained the historical connection write gate")
}
rawConn := &serializedWriteTestConn{}
if err := writeFullToConn(rawConn, []byte("raw")); err != nil {
t.Fatalf("raw write failed: %v", err)
}
if _, ok := transportConnWriteGates.Load(rawConn); ok {
t.Fatal("completed raw write retained a temporary write gate reference")
}
}
+33
View File
@@ -211,6 +211,22 @@ func (c *ClientCommon) setClientTransportProtectionProfile(profile transportProt
c.transportProtection.Store(&profile)
}
func (c *ClientCommon) installInboundTransitionProfile(profile transportProtectionProfile) *transportProtectionProfile {
if c == nil {
return nil
}
next := profile.clone()
c.inboundTransitionProfile.Store(&next)
return &next
}
func (c *ClientCommon) clearInboundTransitionProfile(expected *transportProtectionProfile) {
if c == nil || expected == nil {
return
}
c.inboundTransitionProfile.CompareAndSwap(expected, nil)
}
func (c *ClientCommon) clearClientSecurityProfiles() {
if c == nil {
return
@@ -247,6 +263,7 @@ func (c *ClientCommon) activateClientBootstrapTransportProtection() {
if c == nil || !c.securityConfigured {
return
}
c.inboundTransitionProfile.Store(nil)
c.resetClientNegotiatedSteadyTransportProtection()
c.setClientTransportProtectionProfile(c.securityBootstrap)
}
@@ -406,3 +423,19 @@ func (c *LogicalConn) applyTransportProtectionProfile(profile transportProtectio
state.forwardSecrecyFallback = profile.forwardSecrecyFallback
})
}
func (c *LogicalConn) installInboundTransitionProfile(profile transportProtectionProfile) *transportProtectionProfile {
if c == nil {
return nil
}
next := profile.clone()
c.inboundTransitionProfile.Store(&next)
return &next
}
func (c *LogicalConn) clearInboundTransitionProfile(expected *transportProtectionProfile) {
if c == nil || expected == nil {
return
}
c.inboundTransitionProfile.CompareAndSwap(expected, nil)
}
+111 -2
View File
@@ -4,6 +4,8 @@ import (
"context"
"errors"
"net"
"os"
"sync"
"testing"
"time"
)
@@ -74,7 +76,6 @@ func TestClientSendCtxReturnsContextCanceled(t *testing.T) {
server := newRunningPeerAttachServerForTest(t, func(server *ServerCommon) {
server.SetSecretKey(secret)
})
left, right := net.Pipe()
defer right.Close()
bootstrapPeerAttachConnForTest(t, server, right)
@@ -101,7 +102,6 @@ func TestServerSendCtxReturnsContextCanceled(t *testing.T) {
server := newRunningPeerAttachServerForTest(t, func(server *ServerCommon) {
server.SetSecretKey(secret)
})
left, right := net.Pipe()
defer right.Close()
bootstrapPeerAttachConnForTest(t, server, right)
@@ -133,3 +133,112 @@ func TestServerSendCtxReturnsContextCanceled(t *testing.T) {
t.Fatalf("server SendCtxLogical error = %v, want %v", err, context.Canceled)
}
}
func TestReplyWaitPreservesLegacyDeadlineAndCancelSentinels(t *testing.T) {
client := NewClient().(*ClientCommon)
secret := []byte("0123456789abcdef0123456789abcdef")
client.SetSecretKey(secret)
server := newRunningPeerAttachServerForTest(t, func(server *ServerCommon) {
server.SetSecretKey(secret)
})
clientCancelCtx, clientCancel := context.WithCancel(context.Background())
serverCancelCtx, serverCancel := context.WithCancel(context.Background())
server.SetLink("client-wait-timeout", func(*Message) {})
server.SetLink("client-ctx-cancel", func(*Message) { clientCancel() })
client.SetLink("server-wait-timeout", func(*Message) {})
client.SetLink("server-ctx-cancel", func(*Message) { serverCancel() })
left, right := net.Pipe()
t.Cleanup(func() { _ = right.Close() })
bootstrapPeerAttachConnForTest(t, server, right)
if err := client.ConnectByConn(left); err != nil {
t.Fatalf("client ConnectByConn failed: %v", err)
}
t.Cleanup(func() {
client.setByeFromServer(true)
_ = client.Stop()
})
var logical *LogicalConn
deadline := time.Now().Add(time.Second)
for time.Now().Before(deadline) {
logical = server.GetLogicalConn(client.peerIdentity)
if logical != nil {
break
}
time.Sleep(time.Millisecond)
}
if logical == nil {
t.Fatal("server logical conn not found")
}
if _, err := client.SendWait("client-wait-timeout", []byte("payload"), 50*time.Millisecond); err != os.ErrDeadlineExceeded {
t.Fatalf("client SendWait error=%#v, want exact os.ErrDeadlineExceeded", err)
}
if _, err := client.SendCtx(clientCancelCtx, "client-ctx-cancel", []byte("payload")); err != context.Canceled {
t.Fatalf("client SendCtx error=%#v, want exact context.Canceled", err)
}
if _, err := server.SendWaitLogical(logical, "server-wait-timeout", []byte("payload"), 50*time.Millisecond); err != os.ErrDeadlineExceeded {
t.Fatalf("server SendWaitLogical error=%#v, want exact os.ErrDeadlineExceeded", err)
}
if _, err := server.SendCtxLogical(serverCancelCtx, logical, "server-ctx-cancel", []byte("payload")); err != context.Canceled {
t.Fatalf("server SendCtxLogical error=%#v, want exact context.Canceled", err)
}
}
func TestSetLinkIsSafeDuringConcurrentDispatch(t *testing.T) {
t.Run("client", func(t *testing.T) {
client := NewClient().(*ClientCommon)
testConcurrentHandlerUpdateAndDispatch(
t,
client.SetLink,
client.SetDefaultLink,
client.dispatchMsg,
)
})
t.Run("server", func(t *testing.T) {
server := NewServer().(*ServerCommon)
testConcurrentHandlerUpdateAndDispatch(
t,
server.SetLink,
server.SetDefaultLink,
server.dispatchMsg,
)
})
}
func testConcurrentHandlerUpdateAndDispatch(
t *testing.T,
setLink func(string, func(*Message)),
setDefault func(func(*Message)),
dispatch func(Message),
) {
t.Helper()
const iterations = 1000
start := make(chan struct{})
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
<-start
for i := 0; i < iterations; i++ {
setLink("concurrent-handler", func(*Message) {})
setDefault(func(*Message) {})
}
}()
go func() {
defer wg.Done()
<-start
for i := 0; i < iterations; i++ {
dispatch(Message{TransferMsg: TransferMsg{
Key: "concurrent-handler",
Type: MSG_ASYNC,
}})
}
}()
close(start)
wg.Wait()
}
+5
View File
@@ -22,6 +22,7 @@ type ServerCommon struct {
stopCtx context.Context
maxReadTimeout time.Duration
maxWriteTimeout time.Duration
replyWriteTimeout atomic.Int64
parallelNum int
wg stario.WaitGroup
peerRegistry *serverPeerRegistry
@@ -47,6 +48,7 @@ type ServerCommon struct {
peerAttachAuthRejectCount atomic.Int64
peerAttachDowngradeRejectCount atomic.Int64
peerAttachBindingRejectCount atomic.Int64
linkMu sync.RWMutex
linkFns map[string]func(message *Message)
defaultFns func(message *Message)
noFinSyncMsgMaxKeepSeconds int64
@@ -62,6 +64,8 @@ type ServerCommon struct {
recordRuntime *recordRuntime
bulkRuntime *bulkRuntime
bulkOpenTuning BulkOpenTuning
udpWriteGateOnce sync.Once
udpWriteGate chan struct{}
bulkDedicatedSidecarMu sync.Mutex
bulkDedicatedSidecars map[*LogicalConn]map[uint32]*bulkDedicatedSidecar
connectionRetryState *connectionRetryState
@@ -78,6 +82,7 @@ func NewServer() Server {
server.parallelNum = 0
server.noFinSyncMsgMaxKeepSeconds = 0
server.maxHeartbeatLostSeconds = 300
server.replyWriteTimeout.Store(int64(defaultMessageReplyWriteTimeout))
server.stopCtx, server.stopFn = context.WithCancel(context.Background())
server.SecretKey = nil
server.handshakeRsaKey = defaultRsaKey
+2 -2
View File
@@ -462,8 +462,8 @@ func serverBulkReleaseSender(s *ServerCommon, logical *LogicalConn, transport *T
Chunks: chunks,
}
if transport != nil && transport.IsCurrent() {
return sendBulkReleaseServerTransport(s, transport, req)
return sendBulkReleaseServerTransport(ctx, s, transport, req)
}
return sendBulkReleaseServerLogical(s, logical, req)
return sendBulkReleaseServerLogical(ctx, s, logical, req)
}
}
+22 -3
View File
@@ -1,6 +1,9 @@
package notify
import "context"
import (
"context"
"time"
)
func (s *ServerCommon) DebugMode(dmg bool) {
s.mu.Lock()
@@ -14,6 +17,20 @@ func (s *ServerCommon) IsDebugMode() bool {
return s.debugMode
}
func (s *ServerCommon) SetReplyWriteTimeout(timeout time.Duration) {
if s == nil {
return
}
s.replyWriteTimeout.Store(int64(maxDuration(0, timeout)))
}
func (s *ServerCommon) ReplyWriteTimeout() time.Duration {
if s == nil {
return 0
}
return time.Duration(s.replyWriteTimeout.Load())
}
func (s *ServerCommon) ShowError(std bool) {
s.mu.Lock()
s.showError = std
@@ -57,12 +74,14 @@ func (s *ServerCommon) SetDefaultCommDecode(fn func([]byte, []byte) []byte) {
}
func (s *ServerCommon) SetDefaultLink(fn func(message *Message)) {
s.linkMu.Lock()
defer s.linkMu.Unlock()
s.defaultFns = fn
}
func (s *ServerCommon) SetLink(key string, fn func(*Message)) {
s.mu.Lock()
defer s.mu.Unlock()
s.linkMu.Lock()
defer s.linkMu.Unlock()
s.linkFns[key] = fn
}
+6 -3
View File
@@ -34,12 +34,15 @@ func (s *ServerCommon) dispatchMsg(message Message) {
callFn := func(fn func(*Message)) {
fn(&message)
}
s.linkMu.RLock()
fn, ok := s.linkFns[message.TransferMsg.Key]
if ok {
defaultFn := s.defaultFns
s.linkMu.RUnlock()
if ok && fn != nil {
callFn(fn)
}
if s.defaultFns != nil {
callFn(s.defaultFns)
if defaultFn != nil {
callFn(defaultFn)
}
}
+169 -49
View File
@@ -34,20 +34,35 @@ func (s *ServerCommon) send(c *ClientConn, msg TransferMsg) (WaitMsg, error) {
}
func (s *ServerCommon) sendLogical(logical *LogicalConn, msg TransferMsg) (WaitMsg, error) {
return s.sendLogicalContext(context.Background(), logical, msg, 0)
}
func (s *ServerCommon) sendLogicalContext(ctx context.Context, logical *LogicalConn, msg TransferMsg, writeTimeout time.Duration) (WaitMsg, error) {
if logical == nil {
return s.sendTransport(nil, msg)
return s.sendTransportContextWithWriteTimeout(ctx, nil, msg, writeTimeout)
}
return s.sendTransport(s.resolveOutboundTransport(logical), msg)
return s.sendTransportContextWithWriteTimeout(ctx, s.resolveOutboundTransport(logical), msg, writeTimeout)
}
func (s *ServerCommon) sendTransport(transport *TransportConn, msg TransferMsg) (WaitMsg, error) {
return s.sendTransportContext(context.Background(), transport, msg)
}
func (s *ServerCommon) sendTransportContext(ctx context.Context, transport *TransportConn, msg TransferMsg) (WaitMsg, error) {
return s.sendTransportContextWithWriteTimeout(ctx, transport, msg, 0)
}
func (s *ServerCommon) sendTransportContextWithWriteTimeout(ctx context.Context, transport *TransportConn, msg TransferMsg, writeTimeout time.Duration) (WaitMsg, error) {
if err := s.ensureServerTransportSendReady(transport); err != nil {
return WaitMsg{}, err
}
if s.serverUDPListenerSnapshot() != nil {
return s.sendUDPTransport(transport, msg)
if ctx == nil {
ctx = context.Background()
}
return s.sendTUTransport(transport, msg)
if s.serverUDPListenerSnapshot() != nil {
return s.sendUDPTransportContextWithWriteTimeout(ctx, transport, msg, writeTimeout)
}
return s.sendTUTransportContextWithWriteTimeout(ctx, transport, msg, writeTimeout)
}
func (s *ServerCommon) sendTU(c *ClientConn, msg TransferMsg) (WaitMsg, error) {
@@ -62,6 +77,23 @@ func (s *ServerCommon) sendTULogical(logical *LogicalConn, msg TransferMsg) (Wai
}
func (s *ServerCommon) sendTUTransport(transport *TransportConn, msg TransferMsg) (WaitMsg, error) {
return s.sendTUTransportContext(context.Background(), transport, msg)
}
func (s *ServerCommon) sendTUTransportContext(ctx context.Context, transport *TransportConn, msg TransferMsg) (WaitMsg, error) {
return s.sendTUTransportContextWithWriteTimeout(ctx, transport, msg, 0)
}
func (s *ServerCommon) sendTUTransportContextWithWriteTimeout(ctx context.Context, transport *TransportConn, msg TransferMsg, writeTimeout time.Duration) (WaitMsg, error) {
if err := s.ensureServerTransportSendReady(transport); err != nil {
return WaitMsg{}, err
}
if ctx == nil {
ctx = context.Background()
}
if err := ctx.Err(); err != nil {
return WaitMsg{}, err
}
var wait WaitMsg
if msg.Type != MSG_SYNC_REPLY && msg.Type != MSG_KEY_CHANGE && msg.Type != MSG_SYS_REPLY || msg.ID == 0 {
msg.ID = atomic.AddUint64(&s.msgID, 1)
@@ -74,6 +106,8 @@ func (s *ServerCommon) sendTUTransport(transport *TransportConn, msg TransferMsg
if err != nil {
return WaitMsg{}, err
}
env.controlCtx = ctx
env.controlTimeout = writeTimeout
if requiresSignalReplyWait(msg) {
wait = s.getPendingWaitPool().createAndStoreWithScope(msg, serverTransportScopeForTransport(transport))
}
@@ -121,12 +155,18 @@ func (s *ServerCommon) sendWaitLogical(logical *LogicalConn, msg TransferMsg, ti
}
func (s *ServerCommon) sendTransportWait(transport *TransportConn, msg TransferMsg, timeout time.Duration) (Message, error) {
data, err := s.sendTransport(transport, msg)
ctx := context.Background()
cancel := func() {}
if timeout != 0 {
ctx, cancel = context.WithTimeout(ctx, timeout)
}
defer cancel()
data, err := s.sendTransportContext(ctx, transport, msg)
if err != nil {
return Message{}, err
return Message{}, publicContextSendError(ctx, err)
}
stopCh := sessionStopChan(s.serverStopContextSnapshot())
if timeout.Seconds() == 0 {
if timeout == 0 {
msg, ok := <-data.Reply
if !ok {
return msg, pendingWaitClosedErrorWith(stopCh, transportDetachedErrorForTransport(transport))
@@ -134,7 +174,7 @@ func (s *ServerCommon) sendTransportWait(transport *TransportConn, msg TransferM
return msg, nil
}
select {
case <-time.After(timeout):
case <-ctx.Done():
s.getPendingWaitPool().removeAndClose(data.TransferMsg.ID)
return Message{}, os.ErrDeadlineExceeded
case <-stopCh:
@@ -191,14 +231,14 @@ func (s *ServerCommon) SendCtxTransport(ctx context.Context, t *TransportConn, k
}
func (s *ServerCommon) sendCtxTransport(t *TransportConn, msg TransferMsg, ctx context.Context) (Message, error) {
data, err := s.sendTransport(t, msg)
if err != nil {
return Message{}, err
}
stopCh := sessionStopChan(s.serverStopContextSnapshot())
if ctx == nil {
ctx = context.Background()
}
data, err := s.sendTransportContext(ctx, t, msg)
if err != nil {
return Message{}, publicContextSendError(ctx, err)
}
stopCh := sessionStopChan(s.serverStopContextSnapshot())
select {
case <-ctx.Done():
s.getPendingWaitPool().removeAndClose(data.TransferMsg.ID)
@@ -307,6 +347,23 @@ func (s *ServerCommon) sendUDPLogical(logical *LogicalConn, msg TransferMsg) (Wa
}
func (s *ServerCommon) sendUDPTransport(transport *TransportConn, msg TransferMsg) (WaitMsg, error) {
return s.sendUDPTransportContext(context.Background(), transport, msg)
}
func (s *ServerCommon) sendUDPTransportContext(ctx context.Context, transport *TransportConn, msg TransferMsg) (WaitMsg, error) {
return s.sendUDPTransportContextWithWriteTimeout(ctx, transport, msg, 0)
}
func (s *ServerCommon) sendUDPTransportContextWithWriteTimeout(ctx context.Context, transport *TransportConn, msg TransferMsg, writeTimeout time.Duration) (WaitMsg, error) {
if ctx == nil {
ctx = context.Background()
}
if err := s.ensureServerTransportSendReady(transport); err != nil {
return WaitMsg{}, err
}
if err := ctx.Err(); err != nil {
return WaitMsg{}, err
}
var wait WaitMsg
if msg.Type != MSG_SYNC_REPLY && msg.Type != MSG_KEY_CHANGE && msg.Type != MSG_SYS_REPLY || msg.ID == 0 {
msg.ID = uint64(time.Now().UnixNano()) + rand.Uint64() + rand.Uint64()
@@ -315,6 +372,8 @@ func (s *ServerCommon) sendUDPTransport(transport *TransportConn, msg TransferMs
if err != nil {
return WaitMsg{}, err
}
env.controlCtx = ctx
env.controlTimeout = writeTimeout
if requiresSignalReplyWait(msg) {
wait = s.getPendingWaitPool().createAndStoreWithScope(msg, serverTransportScopeForTransport(transport))
}
@@ -347,14 +406,20 @@ func (s *ServerCommon) sendEnvelopeTransport(transport *TransportConn, env Envel
if logical == nil {
return transportDetachedErrorForTransport(transport)
}
payload, err := s.encodeEnvelopePayloadLogical(logical, env)
var payload []byte
var err error
if env.transportProfile != nil {
payload, err = s.encodeEnvelopePayloadInbound(logical, env, env.transportProfile)
} else {
payload, err = s.encodeEnvelopePayloadLogical(logical, env)
}
if err != nil {
return err
}
if batchedControlEnvelope(env) {
return s.writeControlEnvelopePayload(logical, transport, nil, payload)
return s.writeControlEnvelopePayload(logical, transport, nil, env.controlContext(), payload, env.controlPriority, env.controlTimeout)
}
return s.writeEnvelopePayload(logical, transport, nil, payload)
return s.writeEnvelopePayloadContextTimeout(env.controlContext(), logical, transport, nil, payload, env.controlTimeout)
}
func (s *ServerCommon) sendEnvelopeInboundTransport(logical *LogicalConn, transport *TransportConn, conn net.Conn, env Envelope) error {
@@ -376,34 +441,34 @@ func (s *ServerCommon) sendEnvelopeInboundTransportWithProfile(logical *LogicalC
return err
}
if batchedControlEnvelope(env) {
return s.writeControlEnvelopePayload(logical, transport, conn, payload)
return s.writeControlEnvelopePayload(logical, transport, conn, env.controlContext(), payload, env.controlPriority, env.controlTimeout)
}
return s.writeEnvelopePayload(logical, transport, conn, payload)
return s.writeEnvelopePayloadContextTimeout(env.controlContext(), logical, transport, conn, payload, env.controlTimeout)
}
func (s *ServerCommon) writeControlEnvelopePayload(logical *LogicalConn, transport *TransportConn, conn net.Conn, payload []byte) error {
func (s *ServerCommon) writeControlEnvelopePayload(logical *LogicalConn, transport *TransportConn, conn net.Conn, ctx context.Context, payload []byte, priority controlPriority, writeTimeout time.Duration) error {
if logical == nil {
return transportDetachedErrorForPeer(logical, transport)
}
if s.serverUDPListenerSnapshot() != nil {
return s.writeEnvelopePayload(logical, transport, conn, payload)
return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, writeTimeout)
}
binding := logical.transportBindingSnapshot()
if binding == nil || binding.queueSnapshot() == nil {
return s.writeEnvelopePayload(logical, transport, conn, payload)
return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, writeTimeout)
}
boundConn := binding.connSnapshot()
if boundConn == nil || isPacketTransportConn(boundConn) {
return s.writeEnvelopePayload(logical, transport, conn, payload)
return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, writeTimeout)
}
if conn != nil && conn != boundConn {
return s.writeEnvelopePayload(logical, transport, conn, payload)
return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, writeTimeout)
}
sender := binding.controlBatchSenderSnapshot()
if sender == nil {
return s.writeEnvelopePayload(logical, transport, conn, payload)
return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, writeTimeout)
}
return sender.submit(payload, writeDeadlineFromTimeout(logical.maxWriteTimeoutSnapshot()))
return sender.submitContext(ctx, payload, shorterPositiveDuration(logical.maxWriteTimeoutSnapshot(), writeTimeout), priority)
}
func (s *ServerCommon) encodeEnvelopePayloadInbound(logical *LogicalConn, env Envelope, profile *transportProtectionProfile) ([]byte, error) {
@@ -418,6 +483,10 @@ func (s *ServerCommon) encodeEnvelopePayloadInbound(logical *LogicalConn, env En
}
func (s *ServerCommon) sendTransferInbound(logical *LogicalConn, transport *TransportConn, conn net.Conn, profile *transportProtectionProfile, msg TransferMsg) error {
return s.sendTransferInboundContext(context.Background(), logical, transport, conn, profile, msg, 0)
}
func (s *ServerCommon) sendTransferInboundContext(ctx context.Context, logical *LogicalConn, transport *TransportConn, conn net.Conn, profile *transportProtectionProfile, msg TransferMsg, writeTimeout time.Duration) error {
if logical == nil && transport != nil {
logical = transport.logicalConnSnapshot()
}
@@ -428,10 +497,30 @@ func (s *ServerCommon) sendTransferInbound(logical *LogicalConn, transport *Tran
if err != nil {
return err
}
env.controlCtx = ctx
env.controlTimeout = writeTimeout
return s.sendEnvelopeInboundTransportWithProfile(logical, transport, conn, profile, env)
}
func (s *ServerCommon) writeEnvelopePayload(logical *LogicalConn, transport *TransportConn, conn net.Conn, payload []byte) error {
return s.writeEnvelopePayloadContext(context.Background(), logical, transport, conn, payload)
}
func (s *ServerCommon) writeEnvelopePayloadContext(ctx context.Context, logical *LogicalConn, transport *TransportConn, conn net.Conn, payload []byte) error {
return s.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, payload, 0)
}
func (s *ServerCommon) writeEnvelopePayloadContextTimeout(ctx context.Context, logical *LogicalConn, transport *TransportConn, conn net.Conn, payload []byte, writeTimeout time.Duration) error {
if ctx == nil {
ctx = context.Background()
}
if err := ctx.Err(); err != nil {
return err
}
if logical == nil {
return transportDetachedErrorForPeer(logical, transport)
}
writeTimeout = shorterPositiveDuration(logical.maxWriteTimeoutSnapshot(), writeTimeout)
udpListener := s.serverUDPListenerSnapshot()
queue := s.serverQueueSnapshot()
if queue == nil {
@@ -441,12 +530,18 @@ func (s *ServerCommon) writeEnvelopePayload(logical *LogicalConn, transport *Tra
if transport == nil || transport.RemoteAddr() == nil {
return transportDetachedErrorForTransport(transport)
}
if timeout := logical.maxWriteTimeoutSnapshot(); timeout > 0 {
_ = udpListener.SetWriteDeadline(time.Now().Add(timeout))
}
data := queue.BuildMessage(payload)
_, err := udpListener.WriteTo(data, transport.RemoteAddr())
return err
deadline := earlierWriteDeadline(writeDeadlineFromTimeout(writeTimeout), contextDeadline(ctx))
return s.withUDPWriteLockDeadline(ctx, deadline, func() error {
if !deadline.IsZero() {
if err := udpListener.SetWriteDeadline(deadline); err != nil {
return err
}
defer func() { _ = udpListener.SetWriteDeadline(time.Time{}) }()
}
_, err := udpListener.WriteTo(data, transport.RemoteAddr())
return err
})
}
var binding *transportBinding
if logical != nil {
@@ -456,33 +551,58 @@ func (s *ServerCommon) writeEnvelopePayload(logical *LogicalConn, transport *Tra
if binding == nil {
return os.ErrClosed
}
return binding.withConnWriteLock(func(conn net.Conn) error {
if timeout := logical.maxWriteTimeoutSnapshot(); timeout > 0 {
if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
return err
}
}
lockAcquired, err := binding.withConnWriteLockContextStopTimeout(ctx, nil, writeTimeout, func(conn net.Conn) error {
return writeFramedPayloadUnlocked(conn, queue, payload)
})
if lockAcquired && err != nil {
binding.closeConn()
}
return err
}
if binding != nil && binding.connSnapshot() == conn {
return binding.withConnWriteLock(func(conn net.Conn) error {
if timeout := logical.maxWriteTimeoutSnapshot(); timeout > 0 {
if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
return err
}
}
lockAcquired, err := binding.withConnWriteLockContextStopTimeout(ctx, nil, writeTimeout, func(conn net.Conn) error {
return writeFramedPayloadUnlocked(conn, queue, payload)
})
}
return withRawConnWriteLock(conn, func(conn net.Conn) error {
if timeout := logical.maxWriteTimeoutSnapshot(); timeout > 0 {
if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
return err
}
if lockAcquired && err != nil {
binding.closeConn()
}
return err
}
if err := ctx.Err(); err != nil {
return err
}
deadline := earlierWriteDeadline(writeDeadlineFromTimeout(writeTimeout), contextDeadline(ctx))
writeStarted, err := withRawConnWriteLockContextDeadline(ctx, conn, deadline, func(conn net.Conn) error {
return writeFramedPayloadUnlocked(conn, queue, payload)
})
if writeStarted && err != nil && !isPacketTransportConn(conn) {
_ = conn.Close()
}
return err
}
func (s *ServerCommon) withUDPWriteLock(ctx context.Context, fn func() error) error {
return s.withUDPWriteLockDeadline(ctx, contextDeadline(ctx), fn)
}
func (s *ServerCommon) withUDPWriteLockDeadline(ctx context.Context, deadline time.Time, fn func() error) error {
if s == nil {
return net.ErrClosed
}
if ctx == nil {
ctx = context.Background()
}
s.udpWriteGateOnce.Do(func() {
s.udpWriteGate = newConnWriteGate()
})
if err := lockWriteGateContextDeadline(ctx, nil, s.udpWriteGate, deadline); err != nil {
return err
}
defer func() { s.udpWriteGate <- struct{}{} }()
if err := ctx.Err(); err != nil {
return err
}
return fn()
}
func (s *ServerCommon) dispatchEnvelope(logical *LogicalConn, transport *TransportConn, conn net.Conn, env Envelope, now time.Time) {
+33 -3
View File
@@ -671,6 +671,21 @@ func (s *streamHandle) CloseWrite() error {
return s.close(false)
}
func (s *streamHandle) newControlContext() (context.Context, func(), error) {
if s == nil {
return nil, func() {}, io.ErrClosedPipe
}
s.mu.Lock()
parent := context.Background()
writeTimeout := s.writeTimeout
s.mu.Unlock()
ctx, cancel, _, err := s.newWriteContext(parent, writeTimeout)
if err != nil {
return nil, func() {}, err
}
return ctx, cancel, nil
}
func (s *streamHandle) close(full bool) error {
if s == nil {
return nil
@@ -692,7 +707,12 @@ func (s *streamHandle) close(full bool) error {
s.mu.Unlock()
if closeFn != nil {
if err := closeFn(context.Background(), s, true); err != nil && !errors.Is(err, errStreamNotFound) {
ctx, cancel, err := s.newControlContext()
if err != nil {
return err
}
defer cancel()
if err := closeFn(ctx, s, true); err != nil && !errors.Is(err, errStreamNotFound) {
return err
}
}
@@ -716,7 +736,12 @@ func (s *streamHandle) close(full bool) error {
s.mu.Unlock()
if closeFn != nil {
if err := closeFn(context.Background(), s, full); err != nil && !errors.Is(err, errStreamNotFound) {
ctx, cancel, err := s.newControlContext()
if err != nil {
return err
}
defer cancel()
if err := closeFn(ctx, s, full); err != nil && !errors.Is(err, errStreamNotFound) {
return err
}
}
@@ -758,7 +783,12 @@ func (s *streamHandle) Reset(err error) error {
s.mu.Unlock()
if resetFn != nil {
if sendErr := resetFn(context.Background(), s, streamResetMessage(resetErr)); sendErr != nil {
ctx, cancel, err := s.newControlContext()
if err != nil {
return err
}
defer cancel()
if sendErr := resetFn(ctx, s, streamResetMessage(resetErr)); sendErr != nil {
return sendErr
}
}
+24 -1
View File
@@ -332,18 +332,37 @@ func (s *streamBatchSender) flush(requests []streamBatchRequest) error {
return err
}
writeTimeout := s.transportWriteTimeout()
requestDeadline := streamBatchRequestsEarliestDeadline(requests)
writeCtx := context.Background()
cancelWriteCtx := func() {}
if !requestDeadline.IsZero() {
writeCtx, cancelWriteCtx = context.WithDeadline(writeCtx, requestDeadline)
}
defer cancelWriteCtx()
payloadBytes := 0
for _, payload := range payloads {
payloadBytes += len(payload)
}
started := time.Now()
err = s.binding.withConnWriteLockDeadline(writeDeadlineFromTimeout(writeTimeout), func(conn net.Conn) error {
lockAcquired, err := s.binding.withConnWriteLockContextStopDeadlineManaged(writeCtx, s.stopCh, writeDeadlineFromTimeout(writeTimeout), func(conn net.Conn) error {
return writeFramedPayloadBatchUnlocked(conn, queue, payloads)
})
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.
s.binding.closeConn()
}
return err
}
func streamBatchRequestsEarliestDeadline(requests []streamBatchRequest) time.Time {
var deadline time.Time
for _, req := range requests {
deadline = earlierWriteDeadline(deadline, req.deadline)
}
return deadline
}
func (s *streamBatchSender) transportWriteTimeout() time.Duration {
if s == nil || s.writeTimeoutProvider == nil {
return 0
@@ -508,6 +527,10 @@ func (s *streamBatchSender) stop() {
close(s.stopCh)
})
<-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.
s.flushMu.Lock()
s.flushMu.Unlock()
}
func (s *streamBatchSender) failPending(err error) {
+1 -1
View File
@@ -300,5 +300,5 @@ func (s *ServerCommon) sendFastStreamDataTransport(ctx context.Context, logical
if err != nil {
return err
}
return s.writeEnvelopePayload(logical, transport, nil, payload)
return s.writeEnvelopePayloadContext(ctx, logical, transport, nil, payload)
}
+42
View File
@@ -162,6 +162,48 @@ func TestStreamCloseWriteKeepsReadSideAliveTCP(t *testing.T) {
waitForStreamContextDone(t, stream.Context(), 2*time.Second)
}
func TestStreamCloseAndResetHonorWriteTimeout(t *testing.T) {
newBlockingStream := func(t *testing.T) *streamHandle {
t.Helper()
return newStreamHandle(context.Background(), newStreamRuntime("stream-control-timeout"), clientFileScope(), StreamOpenRequest{
StreamID: "stream-control-timeout",
WriteTimeout: 30 * time.Millisecond,
}, 0, nil, nil, 0,
func(ctx context.Context, _ *streamHandle, _ bool) error {
<-ctx.Done()
return ctx.Err()
},
func(ctx context.Context, _ *streamHandle, _ string) error {
<-ctx.Done()
return ctx.Err()
}, nil, streamConfig{})
}
t.Run("close", func(t *testing.T) {
stream := newBlockingStream(t)
started := time.Now()
err := stream.Close()
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("Close error = %v, want context deadline exceeded", err)
}
if elapsed := time.Since(started); elapsed > time.Second {
t.Fatalf("Close remained blocked for %s", elapsed)
}
})
t.Run("reset", func(t *testing.T) {
stream := newBlockingStream(t)
started := time.Now()
err := stream.Reset(errors.New("reset"))
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("Reset error = %v, want context deadline exceeded", err)
}
if elapsed := time.Since(started); elapsed > time.Second {
t.Fatalf("Reset remained blocked for %s", elapsed)
}
})
}
func TestStreamCloseFullStopsPeerWritesTCP(t *testing.T) {
server := NewServer().(*ServerCommon)
if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil {
+291 -16
View File
@@ -2,18 +2,32 @@ package notify
import (
"b612.me/stario"
"context"
"net"
"sync"
"sync/atomic"
"time"
)
const transportWorkerStopWait = time.Second
// transportBinding models the currently attached physical transport for a
// logical session. The binding can be swapped later without forcing callers to
// reach into raw conn fields directly.
type transportBinding struct {
conn net.Conn
queue *stario.StarQueue
writeMu sync.Mutex
conn net.Conn
queue *stario.StarQueue
writeGateOnce sync.Once
writeGate chan struct{}
writeGateRef *connWriteGateRef
writeGateDone sync.Once
writeActive atomic.Int64
writeStopping atomic.Bool
writeDrainOnce sync.Once
writeDrain chan struct{}
writeDrainDoneOnce sync.Once
adaptiveTx adaptiveTxState
@@ -31,10 +45,7 @@ func newTransportBinding(conn net.Conn, queue *stario.StarQueue) *transportBindi
if conn == nil && queue == nil {
return nil
}
return &transportBinding{
conn: conn,
queue: queue,
}
return &transportBinding{conn: conn, queue: queue}
}
func (b *transportBinding) connSnapshot() net.Conn {
@@ -59,8 +70,11 @@ func (b *transportBinding) withConnWriteLockDeadline(deadline time.Time, fn func
if b == nil {
return net.ErrClosed
}
b.writeMu.Lock()
defer b.writeMu.Unlock()
if err := b.beginConnWrite(); err != nil {
return err
}
<-b.writeGateSnapshot()
defer b.unlockConnWrite()
conn := b.connSnapshot()
if conn == nil {
return net.ErrClosed
@@ -76,6 +90,200 @@ func (b *transportBinding) withConnWriteLockDeadline(deadline time.Time, fn func
return fn(conn)
}
func (b *transportBinding) withConnWriteLockContextTimeout(ctx context.Context, timeout time.Duration, fn func(net.Conn) error) (bool, error) {
return b.withConnWriteLockContextStopTimeout(ctx, nil, timeout, fn)
}
func (b *transportBinding) withConnWriteLockContextStopTimeout(ctx context.Context, stop <-chan struct{}, timeout time.Duration, fn func(net.Conn) error) (bool, error) {
return b.withConnWriteLockContextStopDeadline(ctx, stop, writeDeadlineFromTimeout(timeout), fn)
}
// withConnWriteLockContextStopDeadline carries the caller's context deadline
// into the physical socket write. A context only used for queue admission is
// otherwise unable to interrupt a net.Conn.Write once the write has started.
func (b *transportBinding) withConnWriteLockContextStopDeadline(ctx context.Context, stop <-chan struct{}, deadline time.Time, fn func(net.Conn) error) (bool, error) {
return b.withConnWriteLockContextStopDeadlineMode(ctx, stop, deadline, true, fn)
}
// Sender-owned writes are drained by the sender's stop/flush lifecycle. They
// only need the shutdown admission check, avoiding activity-counter work on
// the bulk, stream, and control hot paths.
func (b *transportBinding) withConnWriteLockContextStopDeadlineManaged(ctx context.Context, stop <-chan struct{}, deadline time.Time, fn func(net.Conn) error) (bool, error) {
return b.withConnWriteLockContextStopDeadlineMode(ctx, stop, deadline, false, fn)
}
func (b *transportBinding) withConnWriteLockContextStopDeadlineMode(ctx context.Context, stop <-chan struct{}, deadline time.Time, trackActivity bool, fn func(net.Conn) error) (bool, error) {
if b == nil {
return false, net.ErrClosed
}
if ctx == nil {
ctx = context.Background()
}
if trackActivity {
if err := b.beginConnWrite(); err != nil {
return false, err
}
} else if b.writeStopping.Load() {
return false, net.ErrClosed
}
deadline = earlierWriteDeadline(deadline, contextDeadline(ctx))
if err := lockWriteGateContextDeadline(ctx, stop, b.writeGateSnapshot(), deadline); err != nil {
if trackActivity {
b.finishConnWrite()
}
return false, err
}
if trackActivity {
defer b.unlockConnWrite()
} else {
defer b.unlockConnWriteManaged()
}
if err := ctx.Err(); err != nil {
return false, err
}
conn := b.connSnapshot()
if conn == nil {
return false, net.ErrClosed
}
if !deadline.IsZero() {
if err := conn.SetWriteDeadline(deadline); err != nil {
return true, err
}
defer func() {
_ = conn.SetWriteDeadline(time.Time{})
}()
}
return true, fn(conn)
}
func contextDeadline(ctx context.Context) time.Time {
if ctx == nil {
return time.Time{}
}
deadline, ok := ctx.Deadline()
if !ok {
return time.Time{}
}
return deadline
}
func earlierWriteDeadline(left time.Time, right time.Time) time.Time {
if left.IsZero() {
return right
}
if right.IsZero() || left.Before(right) {
return left
}
return right
}
func (b *transportBinding) lockConnWriteContext(ctx context.Context) error {
return b.lockConnWriteContextStop(ctx, nil)
}
func (b *transportBinding) lockConnWriteContextStop(ctx context.Context, stop <-chan struct{}) error {
if b == nil {
return net.ErrClosed
}
if ctx == nil {
ctx = context.Background()
}
if err := b.beginConnWrite(); err != nil {
return err
}
if err := lockWriteGateContextDeadline(ctx, stop, b.writeGateSnapshot(), time.Time{}); err != nil {
b.finishConnWrite()
return err
}
return nil
}
func (b *transportBinding) unlockConnWrite() {
if b == nil {
return
}
b.writeGateSnapshot() <- struct{}{}
b.finishConnWrite()
}
func (b *transportBinding) unlockConnWriteManaged() {
if b == nil {
return
}
b.writeGateSnapshot() <- struct{}{}
}
func (b *transportBinding) beginConnWrite() error {
if b == nil || b.writeStopping.Load() {
return net.ErrClosed
}
b.writeActive.Add(1)
if b.writeStopping.Load() {
b.finishConnWrite()
return net.ErrClosed
}
return nil
}
func (b *transportBinding) finishConnWrite() {
if b == nil {
return
}
if b.writeActive.Add(-1) == 0 && b.writeStopping.Load() {
b.signalConnWritesDrained()
}
}
func (b *transportBinding) stopConnWrites() <-chan struct{} {
if b == nil {
done := make(chan struct{})
close(done)
return done
}
b.beginConnWriteShutdown()
done := b.writeDrainSnapshot()
return done
}
func (b *transportBinding) beginConnWriteShutdown() {
if b == nil {
return
}
b.writeStopping.Store(true)
if b.writeActive.Load() == 0 {
b.signalConnWritesDrained()
}
}
func (b *transportBinding) writeDrainSnapshot() chan struct{} {
b.writeDrainOnce.Do(func() {
b.writeDrain = make(chan struct{})
})
return b.writeDrain
}
func (b *transportBinding) signalConnWritesDrained() {
done := b.writeDrainSnapshot()
b.writeDrainDoneOnce.Do(func() {
close(done)
})
}
func (b *transportBinding) writeGateSnapshot() chan struct{} {
if b == nil {
return nil
}
b.writeGateOnce.Do(func() {
if conn := b.connSnapshot(); conn != nil {
b.writeGateRef = retainRawConnWriteGate(conn)
b.writeGate = b.writeGateRef.gate
} else {
b.writeGate = newConnWriteGate()
}
})
return b.writeGate
}
func (b *transportBinding) bulkBatchSenderSnapshotWithCodec(codec bulkBatchCodec, writeTimeout func() time.Duration) *bulkBatchSender {
if b == nil {
return nil
@@ -174,9 +382,29 @@ func (b *transportBinding) serverStreamBatchSenderSnapshot(logical *LogicalConn)
}
func (b *transportBinding) stopBackgroundWorkers() {
b.stopBackgroundWorkersWithClose(false)
}
// stopReplacedTransportBinding interrupts an in-flight physical write before
// waiting for the old binding's workers. A binding may be replaced while
// retaining the same socket (for example, when only its queue changes), so
// that case keeps the connection open.
func stopReplacedTransportBinding(oldBinding *transportBinding, nextBinding *transportBinding, closeConn bool) {
if oldBinding == nil {
return
}
if closeConn && (nextBinding == nil || oldBinding.connSnapshot() != nextBinding.connSnapshot()) {
oldBinding.stopBackgroundWorkersWithClose(true)
return
}
oldBinding.stopBackgroundWorkers()
}
func (b *transportBinding) stopBackgroundWorkersWithClose(closeConn bool) {
if b == nil {
return
}
b.beginConnWriteShutdown()
b.controlMu.Lock()
controlSender := b.controlSender
b.controlMu.Unlock()
@@ -186,13 +414,60 @@ func (b *transportBinding) stopBackgroundWorkers() {
b.bulkMu.Lock()
bulkSender := b.bulkSender
b.bulkMu.Unlock()
if controlSender != nil {
controlSender.stop()
// Closing first is required to interrupt an in-flight physical write during
// actual transport shutdown. Handoff callers normally keep the old socket
// alive, but a bounded wait below closes it if a sender is already inside
// Conn.Write. Reusing a socket after that point could otherwise preserve a
// partial frame and block the handoff forever.
if closeConn {
b.closeConn()
}
if streamSender != nil {
streamSender.stop()
}
if bulkSender != nil {
bulkSender.stop()
workersDone := make(chan struct{})
go func() {
if controlSender != nil {
controlSender.stop()
}
if streamSender != nil {
streamSender.stop()
}
if bulkSender != nil {
bulkSender.stop()
}
b.releaseWriteGate()
close(workersDone)
}()
timer := time.NewTimer(transportWorkerStopWait)
defer timer.Stop()
select {
case <-workersDone:
return
case <-timer.C:
if !closeConn {
// A same-socket handoff may retain the connection only while no old
// sender is active. Once the bounded wait expires the socket is
// unsafe to reuse, so force it closed and let the stopper finish
// asynchronously.
b.closeConn()
}
}
}
func (b *transportBinding) releaseWriteGate() {
if b == nil {
return
}
b.writeGateSnapshot()
<-b.stopConnWrites()
b.writeGateDone.Do(func() {
releaseRawConnWriteGate(b.connSnapshot(), b.writeGateRef)
})
}
func (b *transportBinding) closeConn() {
if b == nil {
return
}
if conn := b.connSnapshot(); conn != nil {
_ = conn.Close()
}
}
+159
View File
@@ -22,6 +22,15 @@ const (
streamAdaptiveSoftPayloadMinSampleBytes = 64 * 1024
streamAdaptiveSoftPayloadGrowSuccesses = 8
controlAdaptiveSoftPayloadMinBytes = 256 * 1024
controlAdaptiveSoftPayloadFallbackBytes = 1 * 1024 * 1024
controlAdaptiveSoftPayloadStartBytes = 4 * 1024 * 1024
controlAdaptiveSoftPayloadMaxBytes = 16 * 1024 * 1024
controlAdaptiveSoftPayloadTargetFlush = 100 * time.Millisecond
controlAdaptiveSoftPayloadSlowFlush = 500 * time.Millisecond
controlAdaptiveSoftPayloadMinSampleBytes = 64 * 1024
controlAdaptiveSoftPayloadGrowSuccesses = 8
streamAdaptiveWaitThresholdMinBytes = 32 * 1024
streamAdaptiveFlushDelayMid = 25 * time.Microsecond
)
@@ -42,6 +51,16 @@ var streamAdaptiveSoftPayloadSteps = [...]int{
streamBatchMaxPayloadBytes,
}
var controlAdaptiveSoftPayloadSteps = [...]int{
256 * 1024,
512 * 1024,
1024 * 1024,
2 * 1024 * 1024,
controlAdaptiveSoftPayloadStartBytes,
8 * 1024 * 1024,
controlAdaptiveSoftPayloadMaxBytes,
}
type adaptiveTxState struct {
mu sync.Mutex
@@ -52,6 +71,24 @@ type adaptiveTxState struct {
streamSoftPayloadBytes int
streamGoodputBytesPerS float64
streamGrowStreak int
controlSoftPayloadBytes int
controlGoodputBytesPerS float64
controlGrowStreak int
}
func (b *transportBinding) controlAdaptiveSoftPayloadBytesSnapshot() int {
if b == nil {
return controlAdaptiveSoftPayloadFallbackBytes
}
return b.adaptiveTx.controlSoftPayloadBytesSnapshot()
}
func (b *transportBinding) observeControlAdaptivePayloadWrite(payloadBytes int, elapsed time.Duration, timeout time.Duration, err error) {
if b == nil {
return
}
b.adaptiveTx.observeControlPayloadWrite(payloadBytes, elapsed, timeout, err)
}
func (b *transportBinding) bulkAdaptiveSoftPayloadBytesSnapshot() int {
@@ -172,6 +209,73 @@ func (s *adaptiveTxState) streamSoftPayloadBytesSnapshot() int {
return s.streamSoftPayloadBytesLocked()
}
func (s *adaptiveTxState) controlSoftPayloadBytesSnapshot() int {
if s == nil {
return controlAdaptiveSoftPayloadStartBytes
}
s.mu.Lock()
defer s.mu.Unlock()
return s.controlSoftPayloadBytesLocked()
}
func (s *adaptiveTxState) observeControlPayloadWrite(payloadBytes int, elapsed time.Duration, timeout time.Duration, err error) {
if s == nil || payloadBytes <= 0 {
return
}
s.mu.Lock()
defer s.mu.Unlock()
current := s.controlSoftPayloadBytesLocked()
target, hasSample := s.observeControlGoodputLocked(payloadBytes, elapsed)
nearTimeout := timeout > 0 && elapsed >= (timeout*3)/4
if isTimeoutLikeError(err) || nearTimeout {
s.controlGrowStreak = 0
if hasSample && target < current {
s.controlSoftPayloadBytes = target
return
}
s.controlSoftPayloadBytes = previousControlAdaptiveSoftPayloadStep(current)
return
}
if err != nil {
s.controlGrowStreak = 0
return
}
if !hasSample {
return
}
if elapsed >= controlAdaptiveSoftPayloadSlowFlush {
s.controlGrowStreak = 0
if target < current {
s.controlSoftPayloadBytes = target
return
}
s.controlSoftPayloadBytes = previousControlAdaptiveSoftPayloadStep(current)
return
}
if target > current {
s.controlGrowStreak++
if s.controlGrowStreak >= controlAdaptiveSoftPayloadGrowSuccesses {
s.controlSoftPayloadBytes = nextControlAdaptiveSoftPayloadStep(current, target)
s.controlGrowStreak = 0
}
return
}
if target < current && elapsed >= controlAdaptiveSoftPayloadTargetFlush*2 {
s.controlSoftPayloadBytes = target
s.controlGrowStreak = 0
return
}
s.controlGrowStreak = 0
}
func (s *adaptiveTxState) controlSoftPayloadBytesLocked() int {
if s.controlSoftPayloadBytes == 0 {
s.controlSoftPayloadBytes = controlAdaptiveSoftPayloadStartBytes
}
return normalizeControlAdaptiveSoftPayloadBytes(s.controlSoftPayloadBytes)
}
func (s *adaptiveTxState) streamWaitThresholdBytesSnapshot() int {
if s == nil {
return streamBatchWaitThreshold
@@ -284,6 +388,24 @@ func (s *adaptiveTxState) observeStreamGoodputLocked(payloadBytes int, elapsed t
return normalizeStreamAdaptiveSoftPayloadBytes(target), true
}
func (s *adaptiveTxState) observeControlGoodputLocked(payloadBytes int, elapsed time.Duration) (int, bool) {
if payloadBytes < controlAdaptiveSoftPayloadMinSampleBytes || elapsed <= 0 {
return 0, false
}
sample := float64(payloadBytes) / elapsed.Seconds()
if sample <= 0 {
return 0, false
}
if s.controlGoodputBytesPerS <= 0 {
s.controlGoodputBytesPerS = sample
} else {
const alpha = 0.25
s.controlGoodputBytesPerS = s.controlGoodputBytesPerS*(1-alpha) + sample*alpha
}
target := int(s.controlGoodputBytesPerS * controlAdaptiveSoftPayloadTargetFlush.Seconds())
return normalizeControlAdaptiveSoftPayloadBytes(target), true
}
func normalizeBulkAdaptiveSoftPayloadBytes(size int) int {
if size <= bulkAdaptiveSoftPayloadMinBytes {
return bulkAdaptiveSoftPayloadMinBytes
@@ -358,6 +480,43 @@ func nextStreamAdaptiveSoftPayloadStep(current int, target int) int {
return streamAdaptiveSoftPayloadStartBytes
}
func normalizeControlAdaptiveSoftPayloadBytes(size int) int {
if size <= controlAdaptiveSoftPayloadMinBytes {
return controlAdaptiveSoftPayloadMinBytes
}
for _, step := range controlAdaptiveSoftPayloadSteps {
if size <= step {
return step
}
}
return controlAdaptiveSoftPayloadMaxBytes
}
func previousControlAdaptiveSoftPayloadStep(current int) int {
current = normalizeControlAdaptiveSoftPayloadBytes(current)
for index := len(controlAdaptiveSoftPayloadSteps) - 1; index >= 0; index-- {
step := controlAdaptiveSoftPayloadSteps[index]
if current > step {
return step
}
}
return controlAdaptiveSoftPayloadMinBytes
}
func nextControlAdaptiveSoftPayloadStep(current int, target int) int {
current = normalizeControlAdaptiveSoftPayloadBytes(current)
target = normalizeControlAdaptiveSoftPayloadBytes(target)
for _, step := range controlAdaptiveSoftPayloadSteps {
if step > current {
if step > target {
return target
}
return step
}
}
return controlAdaptiveSoftPayloadMaxBytes
}
func streamAdaptiveWaitThresholdBytesForSoftPayload(size int) int {
size = normalizeStreamAdaptiveSoftPayloadBytes(size)
threshold := size / 16
+27
View File
@@ -7,6 +7,33 @@ import (
"time"
)
func TestTransportBindingAdaptiveControlStartsLANFriendly(t *testing.T) {
binding := &transportBinding{}
if got, want := binding.controlAdaptiveSoftPayloadBytesSnapshot(), controlAdaptiveSoftPayloadStartBytes; got != want {
t.Fatalf("adaptive control soft payload = %d, want %d", got, want)
}
}
func TestTransportBindingAdaptiveControlShrinksAfterSlowWrite(t *testing.T) {
binding := &transportBinding{}
binding.observeControlAdaptivePayloadWrite(4*1024*1024, 4*time.Second, 30*time.Second, nil)
if got, want := binding.controlAdaptiveSoftPayloadBytesSnapshot(), controlAdaptiveSoftPayloadMinBytes; got != want {
t.Fatalf("adaptive control soft payload = %d, want %d", got, want)
}
}
func TestTransportBindingAdaptiveControlRecoversAfterFastWrites(t *testing.T) {
binding := &transportBinding{}
binding.observeControlAdaptivePayloadWrite(4*1024*1024, 4*time.Second, 30*time.Second, nil)
samples := controlAdaptiveSoftPayloadGrowSuccesses * (len(controlAdaptiveSoftPayloadSteps) - 1)
for i := 0; i < samples; i++ {
binding.observeControlAdaptivePayloadWrite(4*1024*1024, 2*time.Millisecond, 30*time.Second, nil)
}
if got, want := binding.controlAdaptiveSoftPayloadBytesSnapshot(), controlAdaptiveSoftPayloadMaxBytes; got != want {
t.Fatalf("adaptive control soft payload = %d, want %d", got, want)
}
}
func TestTransportBindingAdaptiveBulkSoftPayloadStartsAggressive(t *testing.T) {
binding := &transportBinding{}
if got, want := binding.bulkAdaptiveSoftPayloadBytesSnapshot(), bulkAdaptiveSoftPayloadStartBytes; got != want {
+75
View File
@@ -0,0 +1,75 @@
package notify
import (
"context"
"errors"
"fmt"
)
type TransportSendStage string
const (
TransportSendStageQueue TransportSendStage = "queue"
TransportSendStageWrite TransportSendStage = "write"
TransportSendStageReply TransportSendStage = "reply"
TransportSendStageTransport TransportSendStage = "transport"
)
type TransportSendError struct {
Stage TransportSendStage
Err error
}
func (e *TransportSendError) Error() string {
if e == nil {
return "transport send error"
}
if e.Err == nil {
return fmt.Sprintf("transport send %s failed", e.Stage)
}
return fmt.Sprintf("transport send %s failed: %v", e.Stage, e.Err)
}
func (e *TransportSendError) Unwrap() error {
if e == nil {
return nil
}
return e.Err
}
func TransportSendErrorStage(err error) (TransportSendStage, bool) {
var sendErr *TransportSendError
if !errors.As(err, &sendErr) || sendErr == nil {
return "", false
}
return sendErr.Stage, true
}
func newTransportSendError(stage TransportSendStage, err error) error {
if err == nil {
return nil
}
var sendErr *TransportSendError
if errors.As(err, &sendErr) {
return err
}
return &TransportSendError{Stage: stage, Err: normalizeStreamDeadlineError(err)}
}
// publicContextSendError preserves the legacy public API sentinel when a
// context cancellation is the cause of a transport send failure. Internal
// callers still receive TransportSendError with its stage information.
func publicContextSendError(ctx context.Context, err error) error {
if err == nil || ctx == nil {
return err
}
ctxErr := ctx.Err()
if ctxErr == nil {
return err
}
normalizedCtxErr := normalizeStreamDeadlineError(ctxErr)
if !errors.Is(err, ctxErr) && !errors.Is(err, normalizedCtxErr) {
return err
}
return normalizedCtxErr
}
+124 -14
View File
@@ -2,6 +2,7 @@ package notify
import (
"b612.me/stario"
"context"
"errors"
"io"
"net"
@@ -10,9 +11,15 @@ import (
"time"
)
var transportConnWriteLocks sync.Map
var transportConnWriteGates sync.Map
var errTransportFrameQueueUnavailable = errors.New("transport frame queue is unavailable")
type connWriteGateRef struct {
mu sync.Mutex
gate chan struct{}
refs int
}
type vectoredBuffersWriter interface {
WriteBuffers(*net.Buffers) (int64, error)
}
@@ -124,33 +131,136 @@ func withRawConnWriteLock(conn net.Conn, fn func(net.Conn) error) error {
}
func withRawConnWriteLockDeadline(conn net.Conn, deadline time.Time, fn func(net.Conn) error) error {
_, err := withRawConnWriteLockContextDeadline(context.Background(), conn, deadline, fn)
return err
}
func withRawConnWriteLockContextDeadline(ctx context.Context, conn net.Conn, deadline time.Time, fn func(net.Conn) error) (bool, error) {
if conn == nil {
return net.ErrClosed
return false, net.ErrClosed
}
lock := rawConnWriteLock(conn)
lock.Lock()
defer lock.Unlock()
if ctx == nil {
ctx = context.Background()
}
gateRef := retainRawConnWriteGate(conn)
defer releaseRawConnWriteGate(conn, gateRef)
gate := gateRef.gate
if err := lockWriteGateContextDeadline(ctx, nil, gate, deadline); err != nil {
return false, err
}
defer func() { gate <- struct{}{} }()
if err := ctx.Err(); err != nil {
return false, err
}
deadline = earlierWriteDeadline(deadline, contextDeadline(ctx))
if !deadline.IsZero() {
if err := conn.SetWriteDeadline(deadline); err != nil {
return err
return true, err
}
defer func() {
_ = conn.SetWriteDeadline(time.Time{})
}()
}
return fn(conn)
return true, fn(conn)
}
func rawConnWriteLock(conn net.Conn) *sync.Mutex {
func retainRawConnWriteGate(conn net.Conn) *connWriteGateRef {
if conn == nil {
return &sync.Mutex{}
return &connWriteGateRef{gate: newConnWriteGate(), refs: 1}
}
if lock, ok := transportConnWriteLocks.Load(conn); ok {
return lock.(*sync.Mutex)
for {
candidate := &connWriteGateRef{gate: newConnWriteGate(), refs: 1}
actual, loaded := transportConnWriteGates.LoadOrStore(conn, candidate)
if !loaded {
return candidate
}
ref := actual.(*connWriteGateRef)
ref.mu.Lock()
if ref.refs > 0 {
ref.refs++
ref.mu.Unlock()
return ref
}
ref.mu.Unlock()
transportConnWriteGates.CompareAndDelete(conn, ref)
}
lock := &sync.Mutex{}
actual, _ := transportConnWriteLocks.LoadOrStore(conn, lock)
return actual.(*sync.Mutex)
}
func releaseRawConnWriteGate(conn net.Conn, ref *connWriteGateRef) {
if ref == nil {
return
}
ref.mu.Lock()
if ref.refs > 0 {
ref.refs--
}
remove := ref.refs == 0
ref.mu.Unlock()
if remove && conn != nil {
transportConnWriteGates.CompareAndDelete(conn, ref)
}
}
func newConnWriteGate() chan struct{} {
gate := make(chan struct{}, 1)
gate <- struct{}{}
return gate
}
func lockWriteGateContextDeadline(ctx context.Context, stop <-chan struct{}, gate chan struct{}, deadline time.Time) error {
if ctx == nil {
ctx = context.Background()
}
select {
case <-ctx.Done():
return ctx.Err()
case <-stop:
return net.ErrClosed
default:
}
select {
case <-gate:
return nil
default:
}
if deadline.IsZero() {
select {
case <-ctx.Done():
return ctx.Err()
case <-stop:
return net.ErrClosed
case <-gate:
return nil
}
}
wait := time.Until(deadline)
if wait <= 0 {
return context.DeadlineExceeded
}
timer := time.NewTimer(wait)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-stop:
return net.ErrClosed
case <-timer.C:
return context.DeadlineExceeded
case <-gate:
return nil
}
}
func shorterPositiveDuration(left time.Duration, right time.Duration) time.Duration {
left = maxDuration(0, left)
right = maxDuration(0, right)
if left == 0 {
return right
}
if right == 0 || left < right {
return left
}
return right
}
func writeFramedPayloadUnlocked(conn net.Conn, queue *stario.StarQueue, payload []byte) error {
+1131 -3
View File
File diff suppressed because it is too large Load Diff