fix(notify): 根治低带宽控制面阻塞与传输写入卡死
- 为控制消息增加优先级、公平调度、队列字节预算和自适应批处理 - 支持可取消的写门等待,收紧 shared/dedicated bulk、stream 和 Reply 写入边界 - 修复 bulk reset/close、连接 handoff 和安全 profile 切换时序 - 保留旧取消与超时错误契约,新增阶段化 TransportSendError - 增加 ReplyCtx、ReplyObjCtx、写超时配置及黑洞连接和竞态回归测试
This commit is contained in:
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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) {
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user