diff --git a/bulk.go b/bulk.go index 53a5375..a26f0e6 100644 --- a/bulk.go +++ b/bulk.go @@ -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) diff --git a/bulk_batch_sender.go b/bulk_batch_sender.go index 6146ae8..20f66f4 100644 --- a/bulk_batch_sender.go +++ b/bulk_batch_sender.go @@ -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) { diff --git a/bulk_buffer_release_test.go b/bulk_buffer_release_test.go index 67c098a..64d517c 100644 --- a/bulk_buffer_release_test.go +++ b/bulk_buffer_release_test.go @@ -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") + } +} diff --git a/bulk_control.go b/bulk_control.go index 61ed606..3b366ed 100644 --- a/bulk_control.go +++ b/bulk_control.go @@ -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) { diff --git a/bulk_dedicated.go b/bulk_dedicated.go index b6c172f..ac10163 100644 --- a/bulk_dedicated.go +++ b/bulk_dedicated.go @@ -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) }) } diff --git a/bulk_dedicated_sidecar.go b/bulk_dedicated_sidecar.go index 28d0ad5..3d4627b 100644 --- a/bulk_dedicated_sidecar.go +++ b/bulk_dedicated_sidecar.go @@ -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() + } }) } diff --git a/bulk_fastpath.go b/bulk_fastpath.go index cd3cd61..d635e2d 100644 --- a/bulk_fastpath.go +++ b/bulk_fastpath.go @@ -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: diff --git a/bulk_test.go b/bulk_test.go index 61308e7..deee2c7 100644 --- a/bulk_test.go +++ b/bulk_test.go @@ -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 { diff --git a/client.go b/client.go index 666b325..d437c67 100644 --- a/client.go +++ b/client.go @@ -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 diff --git a/client_bulk.go b/client_bulk.go index 15fbd83..e4d636b 100644 --- a/client_bulk.go +++ b/client_bulk.go @@ -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, diff --git a/client_config.go b/client_config.go index 744ba8a..64c5645 100644 --- a/client_config.go +++ b/client_config.go @@ -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 } diff --git a/client_conn_session.go b/client_conn_session.go index ca3a268..acec1e5 100644 --- a/client_conn_session.go +++ b/client_conn_session.go @@ -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() + } } diff --git a/client_dispatcher.go b/client_dispatcher.go index f1e66c6..e035a0a 100644 --- a/client_dispatcher.go +++ b/client_dispatcher.go @@ -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) } } diff --git a/client_runtime.go b/client_runtime.go index 29ef4ec..28500fe 100644 --- a/client_runtime.go +++ b/client_runtime.go @@ -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 } diff --git a/client_send.go b/client_send.go index 726c51f..ebdf117 100644 --- a/client_send.go +++ b/client_send.go @@ -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) diff --git a/client_session_runtime.go b/client_session_runtime.go index 1c69c79..35dd7da 100644 --- a/client_session_runtime.go +++ b/client_session_runtime.go @@ -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() } diff --git a/client_session_runtime_test.go b/client_session_runtime_test.go index 9933a0b..027ba14 100644 --- a/client_session_runtime_test.go +++ b/client_session_runtime_test.go @@ -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") + } +} diff --git a/client_transport.go b/client_transport.go index b3cf6c6..86c9f52 100644 --- a/client_transport.go +++ b/client_transport.go @@ -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) } diff --git a/control_batch_sender.go b/control_batch_sender.go index 1c1baa1..d81eb97 100644 --- a/control_batch_sender.go +++ b/control_batch_sender.go @@ -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 } diff --git a/envelope.go b/envelope.go index 3e2eb39..12accd8 100644 --- a/envelope.go +++ b/envelope.go @@ -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, } } diff --git a/logical_conn.go b/logical_conn.go index 63aa755..8ec19d1 100644 --- a/logical_conn.go +++ b/logical_conn.go @@ -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()) } diff --git a/msg.go b/msg.go index c77195b..0171bcb 100644 --- a/msg.go +++ b/msg.go @@ -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 { diff --git a/peer_attach_test_helper_test.go b/peer_attach_test_helper_test.go index 5132e55..066f427 100644 --- a/peer_attach_test_helper_test.go +++ b/peer_attach_test_helper_test.go @@ -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 } diff --git a/peer_error.go b/peer_error.go index bb4ef0b..a1733b0 100644 --- a/peer_error.go +++ b/peer_error.go @@ -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 diff --git a/peer_error_test.go b/peer_error_test.go new file mode 100644 index 0000000..b2561ce --- /dev/null +++ b/peer_error_test.go @@ -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) + } + }) + } +} diff --git a/peer_identity.go b/peer_identity.go index 795e877..60e49e4 100644 --- a/peer_identity.go +++ b/peer_identity.go @@ -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 } diff --git a/peer_identity_test.go b/peer_identity_test.go index 154560e..906949b 100644 --- a/peer_identity_test.go +++ b/peer_identity_test.go @@ -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") + } + }) +} diff --git a/release_write_bounds_test.go b/release_write_bounds_test.go new file mode 100644 index 0000000..d90c44d --- /dev/null +++ b/release_write_bounds_test.go @@ -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") + } +} diff --git a/security_profile.go b/security_profile.go index 9b11d9f..051ab26 100644 --- a/security_profile.go +++ b/security_profile.go @@ -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) +} diff --git a/send_state_test.go b/send_state_test.go index 505ceb6..d3aec67 100644 --- a/send_state_test.go +++ b/send_state_test.go @@ -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() +} diff --git a/server.go b/server.go index 11b7651..b16dee9 100644 --- a/server.go +++ b/server.go @@ -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 diff --git a/server_bulk.go b/server_bulk.go index 1a4b512..1b2b196 100644 --- a/server_bulk.go +++ b/server_bulk.go @@ -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) } } diff --git a/server_config.go b/server_config.go index 3e99529..e1adcc8 100644 --- a/server_config.go +++ b/server_config.go @@ -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 } diff --git a/server_dispatcher.go b/server_dispatcher.go index 7c06c8b..2b3847c 100644 --- a/server_dispatcher.go +++ b/server_dispatcher.go @@ -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) } } diff --git a/server_send.go b/server_send.go index 11b94d4..62898cd 100644 --- a/server_send.go +++ b/server_send.go @@ -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) { diff --git a/stream.go b/stream.go index ab05aad..a9b966b 100644 --- a/stream.go +++ b/stream.go @@ -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 } } diff --git a/stream_batch_sender.go b/stream_batch_sender.go index 7a8130d..13a5a3f 100644 --- a/stream_batch_sender.go +++ b/stream_batch_sender.go @@ -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) { diff --git a/stream_fastpath.go b/stream_fastpath.go index 5eb6314..606a18d 100644 --- a/stream_fastpath.go +++ b/stream_fastpath.go @@ -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) } diff --git a/stream_test.go b/stream_test.go index 898be7d..212f162 100644 --- a/stream_test.go +++ b/stream_test.go @@ -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 { diff --git a/transport_binding.go b/transport_binding.go index 6c80a45..916dbbf 100644 --- a/transport_binding.go +++ b/transport_binding.go @@ -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() } } diff --git a/transport_binding_adaptive.go b/transport_binding_adaptive.go index 7562407..cf6a184 100644 --- a/transport_binding_adaptive.go +++ b/transport_binding_adaptive.go @@ -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 diff --git a/transport_binding_adaptive_test.go b/transport_binding_adaptive_test.go index f5ac4e4..3541241 100644 --- a/transport_binding_adaptive_test.go +++ b/transport_binding_adaptive_test.go @@ -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 { diff --git a/transport_send_error.go b/transport_send_error.go new file mode 100644 index 0000000..bd96c92 --- /dev/null +++ b/transport_send_error.go @@ -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 +} diff --git a/transport_write.go b/transport_write.go index 87e80c0..7dd3f04 100644 --- a/transport_write.go +++ b/transport_write.go @@ -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 { diff --git a/transport_write_test.go b/transport_write_test.go index b20c5ce..ecc8cf5 100644 --- a/transport_write_test.go +++ b/transport_write_test.go @@ -5,8 +5,11 @@ import ( "bytes" "context" "errors" + "fmt" "io" "net" + "os" + "strings" "sync" "sync/atomic" "testing" @@ -19,6 +22,85 @@ type serializedWriteTestConn struct { writeCount int32 } +type deadlineRecordingWriteConn struct { + mu sync.Mutex + deadlines []time.Time +} + +func (c *deadlineRecordingWriteConn) Read([]byte) (int, error) { return 0, net.ErrClosed } +func (c *deadlineRecordingWriteConn) Close() error { return nil } +func (c *deadlineRecordingWriteConn) LocalAddr() net.Addr { return nil } +func (c *deadlineRecordingWriteConn) RemoteAddr() net.Addr { return nil } +func (c *deadlineRecordingWriteConn) SetDeadline(time.Time) error { return nil } +func (c *deadlineRecordingWriteConn) SetReadDeadline(time.Time) error { return nil } +func (c *deadlineRecordingWriteConn) Write(p []byte) (int, error) { return len(p), nil } +func (c *deadlineRecordingWriteConn) SetWriteDeadline(deadline time.Time) error { + c.mu.Lock() + c.deadlines = append(c.deadlines, deadline) + c.mu.Unlock() + return nil +} + +func (c *deadlineRecordingWriteConn) deadlineSnapshot() []time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return append([]time.Time(nil), c.deadlines...) +} + +func assertDefaultBulkPhysicalDeadline(t *testing.T, deadlines []time.Time) { + t.Helper() + if len(deadlines) < 2 { + t.Fatalf("physical write deadlines=%v, want bounded deadline followed by reset", deadlines) + } + remaining := time.Until(deadlines[0]) + if remaining < defaultBulkDataWriteTimeout-time.Second || remaining > defaultBulkDataWriteTimeout+time.Second { + t.Fatalf("physical write deadline remaining=%v, want about %v", remaining, defaultBulkDataWriteTimeout) + } + if !deadlines[len(deadlines)-1].IsZero() { + t.Fatalf("physical write deadline was not reset: %v", deadlines) + } +} + +type handoffWriteGateTestConn struct { + firstStarted chan struct{} + secondStarted chan struct{} + releaseFirst chan struct{} + writes atomic.Int32 + active atomic.Int32 + concurrent atomic.Bool +} + +func newHandoffWriteGateTestConn() *handoffWriteGateTestConn { + return &handoffWriteGateTestConn{ + firstStarted: make(chan struct{}), + secondStarted: make(chan struct{}), + releaseFirst: make(chan struct{}), + } +} + +func (c *handoffWriteGateTestConn) Read([]byte) (int, error) { return 0, net.ErrClosed } +func (c *handoffWriteGateTestConn) Close() error { return nil } +func (c *handoffWriteGateTestConn) LocalAddr() net.Addr { return nil } +func (c *handoffWriteGateTestConn) RemoteAddr() net.Addr { return nil } +func (c *handoffWriteGateTestConn) SetDeadline(time.Time) error { return nil } +func (c *handoffWriteGateTestConn) SetReadDeadline(time.Time) error { return nil } +func (c *handoffWriteGateTestConn) SetWriteDeadline(time.Time) error { return nil } +func (c *handoffWriteGateTestConn) Write(p []byte) (int, error) { + active := c.active.Add(1) + if active > 1 { + c.concurrent.Store(true) + } + defer c.active.Add(-1) + switch c.writes.Add(1) { + case 1: + close(c.firstStarted) + <-c.releaseFirst + case 2: + close(c.secondStarted) + } + return len(p), nil +} + func (c *serializedWriteTestConn) Read([]byte) (int, error) { return 0, net.ErrClosed } func (c *serializedWriteTestConn) Close() error { return nil } func (c *serializedWriteTestConn) LocalAddr() net.Addr { return nil } @@ -62,6 +144,156 @@ func TestWriteFullToConnSerializesConcurrentWriters(t *testing.T) { } } +func TestTransportBindingAndRawHandoffSharePhysicalWriteGate(t *testing.T) { + conn := newHandoffWriteGateTestConn() + binding := newTransportBinding(conn, stario.NewQueue()) + bindingDone := make(chan error, 1) + go func() { + bindingDone <- binding.withConnWriteLock(func(conn net.Conn) error { + return writeFullToConnUnlocked(conn, []byte("old-binding-frame")) + }) + }() + select { + case <-conn.firstStarted: + case <-time.After(time.Second): + t.Fatal("binding write did not start") + } + + rawDone := make(chan error, 1) + go func() { + rawDone <- withRawConnWriteLock(conn, func(conn net.Conn) error { + return writeFullToConnUnlocked(conn, []byte("handoff-reply")) + }) + }() + select { + case <-conn.secondStarted: + close(conn.releaseFirst) + <-bindingDone + <-rawDone + t.Fatal("raw handoff write entered Conn.Write before the binding frame completed") + case <-time.After(50 * time.Millisecond): + } + close(conn.releaseFirst) + for name, done := range map[string]<-chan error{"binding": bindingDone, "raw": rawDone} { + select { + case err := <-done: + if err != nil { + t.Fatalf("%s write failed: %v", name, err) + } + case <-time.After(time.Second): + t.Fatalf("%s write did not finish", name) + } + } + if conn.concurrent.Load() { + t.Fatal("binding and raw handoff writes reached Conn.Write concurrently") + } +} + +func TestReleasedTransportBindingKeepsGateUntilActiveWriteFinishes(t *testing.T) { + conn := newHandoffWriteGateTestConn() + oldBinding := newTransportBinding(conn, stario.NewQueue()) + oldDone := make(chan error, 1) + go func() { + oldDone <- oldBinding.withConnWriteLock(func(conn net.Conn) error { + return writeFullToConnUnlocked(conn, []byte("old-frame")) + }) + }() + select { + case <-conn.firstStarted: + case <-time.After(time.Second): + t.Fatal("old binding write did not start") + } + + stopDone := make(chan struct{}) + go func() { + oldBinding.stopBackgroundWorkers() + close(stopDone) + }() + deadline := time.Now().Add(time.Second) + for !oldBinding.writeStopping.Load() && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if !oldBinding.writeStopping.Load() { + close(conn.releaseFirst) + <-oldDone + <-stopDone + t.Fatal("old binding did not enter write shutdown") + } + + newBinding := newTransportBinding(conn, stario.NewQueue()) + newDone := make(chan error, 1) + go func() { + newDone <- newBinding.withConnWriteLock(func(conn net.Conn) error { + return writeFullToConnUnlocked(conn, []byte("new-frame")) + }) + }() + select { + case <-conn.secondStarted: + close(conn.releaseFirst) + <-oldDone + <-newDone + <-stopDone + t.Fatal("replacement binding entered Conn.Write while the released binding still had an active write") + case <-time.After(50 * time.Millisecond): + } + + close(conn.releaseFirst) + for name, done := range map[string]<-chan error{"old binding": oldDone, "new binding": newDone} { + select { + case err := <-done: + if err != nil { + t.Fatalf("%s write failed: %v", name, err) + } + case <-time.After(time.Second): + t.Fatalf("%s write did not finish", name) + } + } + select { + case <-stopDone: + case <-time.After(time.Second): + t.Fatal("old binding shutdown did not finish") + } + newBinding.stopBackgroundWorkers() + if conn.concurrent.Load() { + t.Fatal("released and replacement bindings reached Conn.Write concurrently") + } +} + +func TestRawHandoffWriteDeadlineBoundsBindingGateWait(t *testing.T) { + conn := &serializedWriteTestConn{} + binding := newTransportBinding(conn, stario.NewQueue()) + locked := make(chan struct{}) + release := make(chan struct{}) + bindingDone := make(chan error, 1) + go func() { + bindingDone <- binding.withConnWriteLock(func(net.Conn) error { + close(locked) + <-release + return nil + }) + }() + select { + case <-locked: + case <-time.After(time.Second): + t.Fatal("binding write gate was not acquired") + } + callbackCalled := false + err := withRawConnWriteLockDeadline(conn, time.Now().Add(30*time.Millisecond), func(net.Conn) error { + callbackCalled = true + return nil + }) + close(release) + if bindingErr := <-bindingDone; bindingErr != nil { + t.Fatalf("binding write failed: %v", bindingErr) + } + if !isTimeoutLikeError(err) { + t.Fatalf("raw handoff gate wait error=%v, want timeout", err) + } + if callbackCalled { + t.Fatal("raw handoff callback ran without acquiring the shared binding gate") + } +} + func newTestBulkBatchSender(binding *transportBinding) *bulkBatchSender { return newTestBulkBatchSenderWithWriteTimeout(binding, nil) } @@ -116,6 +348,25 @@ func TestBulkBatchSenderRespectsWriteDeadlineWhenReceiverStalls(t *testing.T) { } } +func TestBulkBatchSenderUsesDefaultPhysicalWriteDeadline(t *testing.T) { + conn := &deadlineRecordingWriteConn{} + sender := newTestBulkBatchSender(newTransportBinding(conn, stario.NewQueue())) + defer sender.stop() + + if err := sender.submitData(context.Background(), 1, 1, bulkFastPathVersionV1, []byte("payload")); err != nil { + t.Fatalf("sender.submitData returned error: %v", err) + } + assertDefaultBulkPhysicalDeadline(t, conn.deadlineSnapshot()) +} + +func TestDedicatedBulkUsesDefaultPhysicalWriteDeadline(t *testing.T) { + conn := &deadlineRecordingWriteConn{} + if err := writeBulkDedicatedRecordWithDeadline(conn, []byte("payload"), time.Time{}); err != nil { + t.Fatalf("writeBulkDedicatedRecordWithDeadline returned error: %v", err) + } + assertDefaultBulkPhysicalDeadline(t, conn.deadlineSnapshot()) +} + func TestStreamBatchSenderRespectsBindingWriteDeadlineWhenReceiverStalls(t *testing.T) { left, right := net.Pipe() defer left.Close() @@ -143,6 +394,230 @@ func TestStreamBatchSenderRespectsBindingWriteDeadlineWhenReceiverStalls(t *test } } +func TestControlBatchSenderCarriesContextDeadlineIntoPhysicalWrite(t *testing.T) { + left, right := net.Pipe() + defer left.Close() + defer right.Close() + + sender := newControlBatchSender(newTransportBinding(left, stario.NewQueue())) + defer sender.stop() + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + err := sender.submitContext(ctx, []byte("deadline-control"), 0, controlPriorityNormal) + if err == nil || !isTimeoutLikeError(err) { + t.Fatalf("control submit error=%v, want timeout-like error", err) + } +} + +func TestBulkBatchSenderCarriesContextDeadlineIntoPhysicalWrite(t *testing.T) { + left, right := net.Pipe() + defer left.Close() + defer right.Close() + + sender := newTestBulkBatchSender(newTransportBinding(left, stario.NewQueue())) + defer sender.stop() + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + err := sender.submitData(ctx, 1, 1, bulkFastPathVersionV1, []byte("deadline-bulk")) + if err == nil || !isTimeoutLikeError(err) { + t.Fatalf("bulk submit error=%v, want timeout-like error", err) + } +} + +func TestStreamBatchSenderCarriesContextDeadlineIntoPhysicalWrite(t *testing.T) { + left, right := net.Pipe() + defer left.Close() + defer right.Close() + + sender := newTestStreamBatchSender(newTransportBinding(left, stario.NewQueue()), nil) + defer sender.stop() + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + err := sender.submitData(ctx, 1, 1, streamFastPathVersionV1, []byte("deadline-stream")) + if err == nil || !isTimeoutLikeError(err) { + t.Fatalf("stream submit error=%v, want timeout-like error", err) + } +} + +func TestTransportBindingStopWithCloseInterruptsPhysicalWrite(t *testing.T) { + left, right := net.Pipe() + defer right.Close() + + binding := newTransportBinding(left, stario.NewQueue()) + sender := newTestBulkBatchSender(binding) + errCh := make(chan error, 1) + go func() { + errCh <- sender.submitData(context.Background(), 1, 1, bulkFastPathVersionV1, []byte("blocked")) + }() + time.Sleep(20 * time.Millisecond) + stopCh := make(chan struct{}) + go func() { + binding.stopBackgroundWorkersWithClose(true) + close(stopCh) + }() + select { + case <-stopCh: + case <-time.After(time.Second): + t.Fatal("stopping transport workers remained blocked on physical write") + } + select { + case <-errCh: + case <-time.After(time.Second): + t.Fatal("blocked submit did not complete after transport close") + } +} + +func TestSameSocketHandoffBoundsBlockedSenderAndClosesSocket(t *testing.T) { + left, right := net.Pipe() + defer right.Close() + + binding := newTransportBinding(left, stario.NewQueue()) + sender := newTestBulkBatchSender(binding) + binding.bulkMu.Lock() + binding.bulkSender = sender + binding.bulkMu.Unlock() + writeDone := make(chan error, 1) + go func() { + writeDone <- sender.submitData(context.Background(), 1, 1, bulkFastPathVersionV1, []byte("handoff-blocked")) + }() + // Give the sender time to leave the queue and enter the physical write. + time.Sleep(20 * time.Millisecond) + + stopDone := make(chan struct{}) + started := time.Now() + go func() { + // The next binding deliberately uses the same socket. The old sender is + // already inside Conn.Write, so retaining the socket is unsafe. + stopReplacedTransportBinding(binding, binding, false) + close(stopDone) + }() + select { + case <-stopDone: + if elapsed := time.Since(started); elapsed > 3*transportWorkerStopWait { + t.Fatalf("same-socket handoff stop took %v, want bounded", elapsed) + } + case <-time.After(3 * transportWorkerStopWait): + t.Fatal("same-socket handoff remained blocked on an in-flight physical write") + } + select { + case err := <-writeDone: + if err == nil { + t.Fatal("blocked sender unexpectedly succeeded after handoff forced socket close") + } + case <-time.After(time.Second): + t.Fatal("blocked sender did not finish after handoff forced socket close") + } +} + +func TestClientCanceledPacketSendDoesNotWrite(t *testing.T) { + conn := newBlockingPacketWriteConn() + client := NewClient().(*ClientCommon) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + client.setClientSessionRuntime(newClientSessionRuntime(conn, stopCtx, stopFn, stario.NewQueue(), 1)) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := client.writeControlPayloadToTransport(ctx, []byte("canceled-packet"), controlPriorityNormal); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled packet send error=%v, want context canceled", err) + } + if got := conn.writeCount.Load(); got != 0 { + t.Fatalf("canceled packet send wrote %d packets", got) + } +} + +func TestServerCanceledUDPSendDoesNotWrite(t *testing.T) { + server := NewServer().(*ServerCommon) + server.alive.Store(true) + stopCtx, stopFn := context.WithCancel(context.Background()) + defer stopFn() + sender, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0}) + if err != nil { + t.Fatalf("ListenUDP sender failed: %v", err) + } + defer sender.Close() + receiver, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0}) + if err != nil { + t.Fatalf("ListenUDP receiver failed: %v", err) + } + defer receiver.Close() + server.setServerSessionRuntime(&serverSessionRuntime{stopCtx: stopCtx, stopFn: stopFn, udpListener: sender}) + logical := server.bootstrapAcceptedLogical("canceled-udp", receiver.LocalAddr(), nil) + if logical == nil || logical.CurrentTransportConn() == nil { + t.Fatal("failed to create UDP transport route") + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := server.sendTransportContext(ctx, logical.CurrentTransportConn(), TransferMsg{Key: "canceled", Type: MSG_ASYNC}); !errors.Is(err, context.Canceled) { + t.Fatalf("canceled UDP send error=%v, want context canceled", err) + } + _ = receiver.SetReadDeadline(time.Now().Add(50 * time.Millisecond)) + var buf [256]byte + if _, _, err := receiver.ReadFromUDP(buf[:]); err == nil { + t.Fatal("canceled UDP send produced a datagram") + } +} + +func TestServerUDPWriteLockSerializesWriters(t *testing.T) { + server := NewServer().(*ServerCommon) + var active atomic.Int32 + var maxActive atomic.Int32 + var first sync.Once + started := make(chan struct{}) + continueFirst := make(chan struct{}) + write := func() error { + current := active.Add(1) + for { + previous := maxActive.Load() + if current <= previous || maxActive.CompareAndSwap(previous, current) { + break + } + } + first.Do(func() { + close(started) + <-continueFirst + }) + active.Add(-1) + return nil + } + firstDone := make(chan error, 1) + secondDone := make(chan error, 1) + go func() { firstDone <- server.withUDPWriteLock(context.Background(), write) }() + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("first UDP writer did not start") + } + go func() { secondDone <- server.withUDPWriteLock(context.Background(), write) }() + time.Sleep(30 * time.Millisecond) + if got := maxActive.Load(); got != 1 { + t.Fatalf("concurrent UDP writers observed = %d, want 1", got) + } + close(continueFirst) + for _, done := range []<-chan error{firstDone, secondDone} { + select { + case err := <-done: + if err != nil { + t.Fatalf("UDP writer returned error: %v", err) + } + case <-time.After(time.Second): + t.Fatal("UDP writer did not finish") + } + } +} + +func TestServerTransportContextSendRejectsMissingRouteBeforeDereference(t *testing.T) { + server := NewServer().(*ServerCommon) + server.alive.Store(true) + msg := TransferMsg{Key: "missing-route", Type: MSG_ASYNC} + + if _, err := server.sendTUTransportContext(context.Background(), nil, msg); !errors.Is(err, errTransportDetached) { + t.Fatalf("missing TU transport error=%v, want transport detached", err) + } + if _, err := server.sendUDPTransportContext(context.Background(), nil, msg); !errors.Is(err, errTransportDetached) { + t.Fatalf("missing UDP transport error=%v, want transport detached", err) + } +} + func TestBulkBatchSenderFlushAggregatesAdaptivePayloadObservation(t *testing.T) { conn := &delayedWriteConn{delay: 20 * time.Millisecond} binding := newTransportBinding(conn, stario.NewQueue()) @@ -186,11 +661,11 @@ func TestControlBatchSenderRespectsWriteDeadlineWhenReceiverStalls(t *testing.T) defer right.Close() sender := newControlBatchSender(newTransportBinding(left, stario.NewQueue())) - deadline := time.Now().Add(50 * time.Millisecond) + writeTimeout := 50 * time.Millisecond errCh := make(chan error, 1) go func() { - errCh <- sender.submit([]byte("payload"), deadline) + errCh <- sender.submit([]byte("payload"), writeTimeout) }() select { @@ -206,6 +681,659 @@ func TestControlBatchSenderRespectsWriteDeadlineWhenReceiverStalls(t *testing.T) } } +func TestControlBatchDirectSubmitDoesNotRequireQueuedState(t *testing.T) { + conn := &serializedWriteTestConn{} + sender := newControlBatchSender(newTransportBinding(conn, stario.NewQueue())) + defer sender.stop() + + submitted, err := sender.tryDirectSubmit(controlBatchRequest{ + ctx: context.Background(), + payload: []byte("direct-control"), + priority: controlPriorityNormal, + queueSize: int64(len("direct-control")), + }) + if err != nil { + t.Fatalf("tryDirectSubmit returned error: %v", err) + } + if !submitted { + t.Fatal("tryDirectSubmit did not use the idle direct path") + } + if got := atomic.LoadInt32(&conn.writeCount); got == 0 { + t.Fatal("direct submit did not write a framed payload") + } + if atomic.LoadInt32(&conn.concurrent) != 0 { + t.Fatal("direct submit performed concurrent transport writes") + } +} + +func TestControlBatchQueuedSubmitAllocatesCompletionState(t *testing.T) { + sender := &controlBatchSender{ + normalCh: make(chan controlBatchRequest, 1), + criticalCh: make(chan controlBatchRequest, 1), + stopCh: make(chan struct{}), + doneCh: make(chan struct{}), + budgetWake: make(chan struct{}, 1), + } + sender.flushMu.Lock() + defer sender.flushMu.Unlock() + + errCh := make(chan error, 1) + go func() { + errCh <- sender.submitContext(context.Background(), []byte("queued-control"), 0, controlPriorityNormal) + }() + + var req controlBatchRequest + select { + case req = <-sender.normalCh: + case <-time.After(time.Second): + t.Fatal("control request did not enter the queue") + } + if req.done == nil { + t.Fatal("queued control request has no completion channel") + } + if req.state == nil { + t.Fatal("queued control request has no cancellation state") + } + sender.finishRequest(req, nil) + + select { + case err := <-errCh: + if err != nil { + t.Fatalf("queued submit returned error: %v", err) + } + case <-time.After(time.Second): + t.Fatal("queued submit did not return after completion") + } +} + +type blockingRecordingControlConn struct { + startOnce sync.Once + unblockOnce sync.Once + startCh chan struct{} + unblockCh chan struct{} + + mu sync.Mutex + data bytes.Buffer +} + +func newBlockingRecordingControlConn() *blockingRecordingControlConn { + return &blockingRecordingControlConn{ + startCh: make(chan struct{}), + unblockCh: make(chan struct{}), + } +} + +func (c *blockingRecordingControlConn) Read([]byte) (int, error) { return 0, net.ErrClosed } +func (c *blockingRecordingControlConn) Close() error { return nil } +func (c *blockingRecordingControlConn) LocalAddr() net.Addr { return nil } +func (c *blockingRecordingControlConn) RemoteAddr() net.Addr { return nil } +func (c *blockingRecordingControlConn) SetDeadline(time.Time) error { + return nil +} +func (c *blockingRecordingControlConn) SetReadDeadline(time.Time) error { + return nil +} +func (c *blockingRecordingControlConn) SetWriteDeadline(time.Time) error { + return nil +} + +func (c *blockingRecordingControlConn) Write(p []byte) (int, error) { + c.startOnce.Do(func() { + close(c.startCh) + <-c.unblockCh + }) + c.mu.Lock() + defer c.mu.Unlock() + return c.data.Write(p) +} + +func (c *blockingRecordingControlConn) bytes() []byte { + c.mu.Lock() + defer c.mu.Unlock() + return append([]byte(nil), c.data.Bytes()...) +} + +func (c *blockingRecordingControlConn) unblock() { + c.unblockOnce.Do(func() { + close(c.unblockCh) + }) +} + +func TestControlBatchSenderPrioritizesCriticalRequestAfterCurrentFlush(t *testing.T) { + conn := newBlockingRecordingControlConn() + sender := newControlBatchSender(newTransportBinding(conn, stario.NewQueue())) + defer sender.stop() + defer conn.unblock() + + firstDone := make(chan error, 1) + go func() { + firstDone <- sender.submitContext(context.Background(), []byte("control-first"), 0, controlPriorityNormal) + }() + select { + case <-conn.startCh: + case <-time.After(time.Second): + t.Fatal("first control write did not start") + } + + normalDone := make(chan error, 1) + go func() { + normalDone <- sender.submitContext(context.Background(), []byte("control-normal"), 0, controlPriorityNormal) + }() + criticalDone := make(chan error, 1) + go func() { + criticalDone <- sender.submitContext(context.Background(), []byte("control-critical"), 0, controlPriorityCritical) + }() + time.Sleep(20 * time.Millisecond) + conn.unblock() + + for name, done := range map[string]<-chan error{ + "first": firstDone, "normal": normalDone, "critical": criticalDone, + } { + select { + case err := <-done: + if err != nil { + t.Fatalf("%s submit failed: %v", name, err) + } + case <-time.After(time.Second): + t.Fatalf("%s submit did not finish", name) + } + } + + written := conn.bytes() + criticalAt := bytes.Index(written, []byte("control-critical")) + normalAt := bytes.Index(written, []byte("control-normal")) + if criticalAt < 0 || normalAt < 0 { + t.Fatalf("missing control payloads in write: critical=%d normal=%d", criticalAt, normalAt) + } + if criticalAt > normalAt { + t.Fatalf("critical control payload was written after normal payload: critical=%d normal=%d", criticalAt, normalAt) + } +} + +func TestControlBatchSenderCancelsQueuedRequestWithoutStoppingSender(t *testing.T) { + conn := newBlockingRecordingControlConn() + sender := newControlBatchSender(newTransportBinding(conn, stario.NewQueue())) + defer sender.stop() + defer conn.unblock() + + firstDone := make(chan error, 1) + go func() { + firstDone <- sender.submitContext(context.Background(), []byte("control-first"), 0, controlPriorityNormal) + }() + select { + case <-conn.startCh: + case <-time.After(time.Second): + t.Fatal("first control write did not start") + } + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + started := time.Now() + err := sender.submitContext(ctx, []byte("control-canceled"), 0, controlPriorityNormal) + if !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("queued submit error = %v, want legacy deadline match", err) + } + if time.Since(started) > 250*time.Millisecond { + t.Fatalf("queued cancellation took too long: %s", time.Since(started)) + } + + conn.unblock() + if err := <-firstDone; err != nil { + t.Fatalf("first submit failed: %v", err) + } + if err := sender.submitContext(context.Background(), []byte("control-after-cancel"), 0, controlPriorityNormal); err != nil { + t.Fatalf("sender stopped after queued cancellation: %v", err) + } + if bytes.Contains(conn.bytes(), []byte("control-canceled")) { + t.Fatal("canceled queued control payload was written") + } +} + +func TestControlBatchSenderCancelsWhileWaitingForSharedWriteLock(t *testing.T) { + binding := newTransportBinding(&serializedWriteTestConn{}, stario.NewQueue()) + sender := newControlBatchSender(binding) + defer sender.stop() + + lockHeld := make(chan struct{}) + releaseLock := make(chan struct{}) + lockDone := make(chan error, 1) + go func() { + lockDone <- binding.withConnWriteLock(func(net.Conn) error { + close(lockHeld) + <-releaseLock + return nil + }) + }() + select { + case <-lockHeld: + case <-time.After(time.Second): + t.Fatal("shared transport write lock was not acquired") + } + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + errCh := make(chan error, 1) + go func() { + errCh <- sender.submitContext(ctx, []byte("control-lock-wait"), 0, controlPriorityNormal) + }() + + var err error + select { + case err = <-errCh: + case <-time.After(250 * time.Millisecond): + close(releaseLock) + <-lockDone + t.Fatal("control submit ignored context while waiting for the shared write lock") + } + stage, ok := TransportSendErrorStage(err) + if !ok || stage != TransportSendStageQueue || !errors.Is(err, os.ErrDeadlineExceeded) { + close(releaseLock) + <-lockDone + t.Fatalf("control lock-wait error=%v stage=%q,%v; want queue deadline", err, stage, ok) + } + if sender.errSnapshot() != nil { + close(releaseLock) + <-lockDone + t.Fatalf("caller cancellation stopped a healthy control sender: %v", sender.errSnapshot()) + } + + close(releaseLock) + if err := <-lockDone; err != nil { + t.Fatalf("shared transport lock holder returned error: %v", err) + } + if err := sender.submitContext(context.Background(), []byte("control-after-lock-cancel"), 0, controlPriorityNormal); err != nil { + t.Fatalf("sender failed after lock-wait cancellation: %v", err) + } +} + +func TestControlBatchSenderIsNotStarvedByContinuousSharedWriters(t *testing.T) { + binding := newTransportBinding(&serializedWriteTestConn{}, stario.NewQueue()) + sender := newControlBatchSender(binding) + defer sender.stop() + + if err := binding.lockConnWriteContext(context.Background()); err != nil { + t.Fatalf("initial shared transport lock failed: %v", err) + } + + ctx, cancel := context.WithTimeout(context.Background(), 150*time.Millisecond) + defer cancel() + controlDone := make(chan error, 1) + go func() { + controlDone <- sender.submitContext(ctx, []byte("control-fairness"), 0, controlPriorityCritical) + }() + + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + if sender.flushMu.TryLock() { + sender.flushMu.Unlock() + time.Sleep(time.Millisecond) + continue + } + break + } + if sender.flushMu.TryLock() { + sender.flushMu.Unlock() + binding.unlockConnWrite() + t.Fatal("control request did not enter the shared write wait") + } + + var writers sync.WaitGroup + for i := 0; i < 64; i++ { + writers.Add(1) + go func() { + defer writers.Done() + _ = binding.withConnWriteLock(func(net.Conn) error { + time.Sleep(3 * time.Millisecond) + return nil + }) + }() + } + time.Sleep(10 * time.Millisecond) + binding.unlockConnWrite() + + err := <-controlDone + writers.Wait() + if err != nil { + t.Fatalf("control request was starved by shared writers: %v", err) + } +} + +func TestControlBatchStartedWriteCancellationReturnsBeforePhysicalWrite(t *testing.T) { + conn := newBlockingRecordingControlConn() + sender := newControlBatchSender(newTransportBinding(conn, stario.NewQueue())) + defer sender.stop() + defer conn.unblock() + + ctx, cancel := context.WithCancel(context.Background()) + payload := []byte("control-started-owned") + done := make(chan error, 1) + go func() { + done <- sender.submitContext(ctx, payload, 0, controlPriorityCritical) + }() + select { + case <-conn.startCh: + case <-time.After(time.Second): + t.Fatal("control write did not start") + } + + cancel() + select { + case err := <-done: + if !errors.Is(err, context.Canceled) { + t.Fatalf("started control cancellation error=%v, want context canceled", err) + } + case <-time.After(150 * time.Millisecond): + t.Fatal("started control write kept the caller blocked after cancellation") + } + for i := range payload { + payload[i] = 'x' + } + if sender.queued.Load() == 0 || sender.queuedSize.Load() == 0 { + t.Fatal("started write released queue ownership before the physical write completed") + } + + conn.unblock() + deadline := time.Now().Add(time.Second) + for (sender.queued.Load() != 0 || sender.queuedSize.Load() != 0) && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if got := sender.queued.Load(); got != 0 { + t.Fatalf("queued requests after physical write = %d, want 0", got) + } + if got := sender.queuedSize.Load(); got != 0 { + t.Fatalf("queued bytes after physical write = %d, want 0", got) + } + if !bytes.Contains(conn.bytes(), []byte("control-started-owned")) { + t.Fatal("sender did not retain the started payload after caller cancellation") + } + if err := sender.submit([]byte("control-after-started-cancel"), 0); err != nil { + t.Fatalf("sender failed after a started caller cancellation: %v", err) + } +} + +func TestControlBatchSenderStopCancelsDirectWaitForSharedWriteLock(t *testing.T) { + binding := newTransportBinding(&serializedWriteTestConn{}, stario.NewQueue()) + sender := newControlBatchSender(binding) + lockHeld := make(chan struct{}) + releaseLock := make(chan struct{}) + var releaseOnce sync.Once + release := func() { releaseOnce.Do(func() { close(releaseLock) }) } + defer sender.stop() + defer release() + + lockDone := make(chan error, 1) + go func() { + lockDone <- binding.withConnWriteLock(func(net.Conn) error { + close(lockHeld) + <-releaseLock + return nil + }) + }() + select { + case <-lockHeld: + case <-time.After(time.Second): + t.Fatal("shared transport write lock was not acquired") + } + + submitDone := make(chan error, 1) + go func() { + submitDone <- sender.submitContext(context.Background(), []byte("control-stop-lock-wait"), 0, controlPriorityNormal) + }() + deadline := time.Now().Add(time.Second) + for time.Now().Before(deadline) { + if sender.flushMu.TryLock() { + sender.flushMu.Unlock() + time.Sleep(time.Millisecond) + continue + } + break + } + if sender.flushMu.TryLock() { + sender.flushMu.Unlock() + t.Fatal("direct control request did not enter the shared-lock wait") + } + + stopDone := make(chan struct{}) + go func() { + sender.stop() + close(stopDone) + }() + select { + case <-stopDone: + case <-time.After(250 * time.Millisecond): + release() + <-lockDone + <-submitDone + <-stopDone + t.Fatal("sender stop did not cancel a direct request waiting for the shared write lock") + } + select { + case err := <-submitDone: + if !errors.Is(err, errTransportDetached) { + t.Fatalf("direct submit after sender stop error=%v, want transport detached", err) + } + case <-time.After(time.Second): + t.Fatal("direct submit did not return after sender stop") + } + + release() + if err := <-lockDone; err != nil { + t.Fatalf("shared transport lock holder returned error: %v", err) + } +} + +func TestControlBatchSenderRejectsOversizedPayloadBeforeWrite(t *testing.T) { + conn := &serializedWriteTestConn{} + sender := newControlBatchSender(newTransportBinding(conn, stario.NewQueue())) + defer sender.stop() + + err := sender.submitContext( + context.Background(), + make([]byte, controlBatchMaxPayloadBytes+1), + 0, + controlPriorityNormal, + ) + stage, ok := TransportSendErrorStage(err) + if !ok || stage != TransportSendStageQueue { + t.Fatalf("oversized submit error=%v stage=%q,%v; want queue-stage rejection", err, stage, ok) + } + if err == nil || !strings.Contains(err.Error(), "stream or bulk") { + t.Fatalf("oversized submit error=%v; want stream or bulk guidance", err) + } + if got := atomic.LoadInt32(&conn.writeCount); got != 0 { + t.Fatalf("oversized control payload performed %d physical writes, want 0", got) + } + if sender.errSnapshot() != nil { + t.Fatalf("oversized caller payload stopped a healthy sender: %v", sender.errSnapshot()) + } + if err := sender.submit([]byte("control-after-oversized"), 0); err != nil { + t.Fatalf("sender failed after oversized payload rejection: %v", err) + } +} + +func TestControlBatchSenderBoundsCriticalBurstWhenNormalWaits(t *testing.T) { + conn := newBlockingRecordingControlConn() + conn.unblock() + sender := newControlBatchSender(newTransportBinding(conn, stario.NewQueue())) + sender.flushMu.Lock() + flushLocked := true + defer func() { + if flushLocked { + sender.flushMu.Unlock() + } + sender.stop() + }() + + firstCriticalDone := make(chan error, 1) + go func() { + firstCriticalDone <- sender.submitContext(context.Background(), []byte("control-critical-first"), 0, controlPriorityCritical) + }() + deadline := time.Now().Add(time.Second) + for sender.queued.Load() != 1 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if sender.queued.Load() != 1 { + t.Fatal("first critical request did not enter the queued path") + } + // Let run consume this request as a one-item batch before adding the rest. + time.Sleep(20 * time.Millisecond) + + normalDone := make(chan error, 1) + go func() { + normalDone <- sender.submitContext(context.Background(), []byte("control-normal-waiting"), 0, controlPriorityNormal) + }() + criticalDone := make([]chan error, controlBatchMaxCriticalBurst) + for i := range criticalDone { + criticalDone[i] = make(chan error, 1) + payload := []byte(fmt.Sprintf("control-critical-%02d", i)) + go func(done chan<- error, payload []byte) { + done <- sender.submitContext(context.Background(), payload, 0, controlPriorityCritical) + }(criticalDone[i], payload) + } + deadline = time.Now().Add(time.Second) + for sender.queued.Load() != int64(controlBatchMaxCriticalBurst+2) && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if got, want := sender.queued.Load(), int64(controlBatchMaxCriticalBurst+2); got != want { + t.Fatalf("queued requests=%d, want %d before releasing the sender", got, want) + } + + sender.flushMu.Unlock() + flushLocked = false + for name, done := range map[string]<-chan error{ + "first critical": firstCriticalDone, + "normal": normalDone, + } { + select { + case err := <-done: + if err != nil { + t.Fatalf("%s submit failed: %v", name, err) + } + case <-time.After(2 * time.Second): + t.Fatalf("%s submit did not finish", name) + } + } + for i, done := range criticalDone { + select { + case err := <-done: + if err != nil { + t.Fatalf("critical submit %d failed: %v", i, err) + } + case <-time.After(2 * time.Second): + t.Fatalf("critical submit %d did not finish", i) + } + } + + written := conn.bytes() + normalAt := bytes.Index(written, []byte("control-normal-waiting")) + if normalAt < 0 { + t.Fatal("normal control payload was not written") + } + criticalBeforeNormal := bytes.Count(written[:normalAt], []byte("control-critical-")) + if criticalBeforeNormal > controlBatchMaxCriticalBurst { + t.Fatalf("normal request waited behind %d critical messages, max burst is %d", criticalBeforeNormal, controlBatchMaxCriticalBurst) + } +} + +func TestControlBatchByteBudgetDoesNotAppendPastSoftLimit(t *testing.T) { + limit := 4 * 1024 * 1024 + if !controlBatchCanAppend(1024*1024, 2*1024*1024, limit) { + t.Fatal("batch should accept request below the cumulative byte limit") + } + if controlBatchCanAppend(3*1024*1024, 2*1024*1024, limit) { + t.Fatal("batch accepted request past the cumulative byte limit") + } + if !controlBatchCanAppend(0, 8*1024*1024, limit) { + t.Fatal("one oversized request must remain sendable as its own batch") + } +} + +func TestControlBatchSubmitReturnsWhenSenderStopsAfterQueue(t *testing.T) { + sender := &controlBatchSender{ + normalCh: make(chan controlBatchRequest, 1), + criticalCh: make(chan controlBatchRequest, 1), + stopCh: make(chan struct{}), + doneCh: make(chan struct{}), + budgetWake: make(chan struct{}, 1), + } + sender.flushMu.Lock() + defer sender.flushMu.Unlock() + + errCh := make(chan error, 1) + go func() { + errCh <- sender.submitContext(context.Background(), []byte("queued-before-stop"), 0, controlPriorityNormal) + }() + deadline := time.Now().Add(time.Second) + for len(sender.normalCh) == 0 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if len(sender.normalCh) == 0 { + t.Fatal("control request did not enter the queue") + } + sender.markFailed(newTransportSendError(TransportSendStageTransport, errTransportDetached)) + + select { + case err := <-errCh: + if err == nil { + t.Fatal("queued submit returned nil after sender stop") + } + stage, ok := TransportSendErrorStage(err) + if !ok || stage != TransportSendStageTransport { + t.Fatalf("queued submit error stage=%q, %v; want %q, true", stage, ok, TransportSendStageTransport) + } + case <-time.After(time.Second): + t.Fatal("queued submit remained blocked after sender stop") + } +} + +func TestControlBatchSenderStopAdmissionReleasesAllQueueBudget(t *testing.T) { + for iteration := 0; iteration < 100; iteration++ { + sender := newControlBatchSender(newTransportBinding(&serializedWriteTestConn{}, stario.NewQueue())) + start := make(chan struct{}) + results := make(chan error, 32) + for i := 0; i < cap(results); i++ { + go func() { + <-start + results <- sender.submit([]byte("stop-admission-race"), 0) + }() + } + close(start) + sender.markFailed(newTransportSendError(TransportSendStageTransport, errTransportDetached)) + sender.stop() + for i := 0; i < cap(results); i++ { + <-results + } + if got := sender.queued.Load(); got != 0 { + t.Fatalf("iteration %d queued requests=%d after stop, want 0", iteration, got) + } + if got := sender.queuedSize.Load(); got != 0 { + t.Fatalf("iteration %d queued bytes=%d after stop, want 0", iteration, got) + } + } +} + +func TestTransportSendErrorPreservesStageAndCause(t *testing.T) { + err := newTransportSendError(TransportSendStageQueue, context.DeadlineExceeded) + if !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("transport send error does not preserve legacy deadline matching: %v", err) + } + stage, ok := TransportSendErrorStage(err) + if !ok || stage != TransportSendStageQueue { + t.Fatalf("transport send stage = %q, %v; want %q, true", stage, ok, TransportSendStageQueue) + } +} + +func TestPublicContextSendErrorRestoresNormalizedDeadlineSentinel(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), time.Millisecond) + defer cancel() + <-ctx.Done() + + err := newTransportSendError(TransportSendStageQueue, ctx.Err()) + if got := publicContextSendError(ctx, err); got != os.ErrDeadlineExceeded { + t.Fatalf("public context send error=%#v, want exact os.ErrDeadlineExceeded", got) + } +} + type blockingPacketWriteConn struct { startCh chan struct{} unblockCh chan struct{} @@ -495,7 +1623,7 @@ func TestTransportBindingStopBackgroundWorkersStopsControlSender(t *testing.T) { sender := binding.controlBatchSenderSnapshot() binding.stopBackgroundWorkers() - err := sender.submit([]byte("payload"), time.Time{}) + err := sender.submit([]byte("payload"), 0) if !errors.Is(err, errTransportDetached) { t.Fatalf("sender.submit after stop = %v, want %v", err, errTransportDetached) }