diff --git a/docs/BENCHMARK.md b/docs/BENCHMARK.md index 9c14bcd..8ce4e99 100644 --- a/docs/BENCHMARK.md +++ b/docs/BENCHMARK.md @@ -61,6 +61,17 @@ one-byte completion acknowledgement, not tunnel setup. In particular, a short flow ratio is a warm-payload diagnostic rather than a connection-establishment benchmark. +In source builds containing the cancellation fix (not the published v1.0.1 +binary), stopping `bench-server` closes its listener and accepted transfers, +including a connection waiting for a transfer slot, and joins the transfer +workers. It no longer waits only for a stalled peer's two-minute transfer +deadline. Stopping `bench-client` interrupts its current owned connection even +while writing the request, moving payload, or waiting for the acknowledgement. +The shared tunnel dialer is not closed by an individual canceled transfer. +An incomplete iteration fails without emitting a successful measurement or +partial summary. Successful measurements retain the same wire format, timing +window, payload limits, and statistics. + ## Namespace topology `scripts/netem-integration.sh` creates two isolated network namespaces joined by diff --git a/internal/netbench/cancellation_test.go b/internal/netbench/cancellation_test.go new file mode 100644 index 0000000..658e919 --- /dev/null +++ b/internal/netbench/cancellation_test.go @@ -0,0 +1,392 @@ +package netbench + +import ( + "context" + "encoding/binary" + "errors" + "fmt" + "io" + "net" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +const ( + benchmarkCancelReturnBudget = 400 * time.Millisecond + benchmarkCancelJoinBudget = 2 * time.Second +) + +// These are actual socket reads. Observing entry does not create an I/O gate +// or consume bytes independently of the production handler. +type benchmarkCancelServerConn struct { + net.Conn + readStarted chan struct{} + uploadStarted chan struct{} + readOnce sync.Once + uploadOnce sync.Once + bytesRead atomic.Int64 +} + +func (c *benchmarkCancelServerConn) Read(p []byte) (int, error) { + c.readOnce.Do(func() { close(c.readStarted) }) + if c.bytesRead.Load() >= headerSize { + c.uploadOnce.Do(func() { close(c.uploadStarted) }) + } + n, err := c.Conn.Read(p) + c.bytesRead.Add(int64(n)) + return n, err +} + +type benchmarkCancelListener struct { + net.Listener + accepted chan *benchmarkCancelServerConn + mu sync.Mutex + closed bool + conns []*benchmarkCancelServerConn +} + +func (l *benchmarkCancelListener) Accept() (net.Conn, error) { + raw, err := l.Listener.Accept() + if err != nil { + return nil, err + } + conn := &benchmarkCancelServerConn{ + Conn: raw, readStarted: make(chan struct{}), uploadStarted: make(chan struct{}), + } + l.mu.Lock() + if l.closed { + l.mu.Unlock() + _ = conn.Close() + return nil, net.ErrClosed + } + l.conns = append(l.conns, conn) + l.mu.Unlock() + // The tests create at most two peers. A bounded notification never + // changes delivery timing or prevents production Accept from returning. + select { + case l.accepted <- conn: + default: + } + return conn, nil +} + +func (l *benchmarkCancelListener) forceCloseAccepted() { + l.mu.Lock() + l.closed = true + conns := append([]*benchmarkCancelServerConn(nil), l.conns...) + l.mu.Unlock() + for _, conn := range conns { + _ = conn.Close() + } +} + +type benchmarkCancelServerFixture struct { + listener *benchmarkCancelListener + cancel context.CancelFunc + done chan struct{} + result chan error + peers []net.Conn +} + +func benchmarkCancelStartServer(t *testing.T) *benchmarkCancelServerFixture { + t.Helper() + raw, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + f := &benchmarkCancelServerFixture{ + listener: &benchmarkCancelListener{Listener: raw, accepted: make(chan *benchmarkCancelServerConn, 4)}, + cancel: cancel, done: make(chan struct{}), result: make(chan error, 1), + } + t.Cleanup(func() { + // Cleanup remains independent of the implementation's cancellation. + cancel() + _ = raw.Close() + f.listener.forceCloseAccepted() + for _, peer := range f.peers { + _ = peer.Close() + } + benchmarkCancelJoin(t, f.done, "server cleanup") + }) + go func() { + defer close(f.done) + // Keep the actual default two-minute handler timeout: the test must + // prove cancellation, not expiration of a shortened socket deadline. + f.result <- (&Server{MaxBytes: 1 << 20, MaxConnections: 1}).Serve(ctx, f.listener) + }() + return f +} + +func (f *benchmarkCancelServerFixture) connect(t *testing.T) (net.Conn, *benchmarkCancelServerConn) { + t.Helper() + peer, err := (&net.Dialer{Timeout: time.Second}).DialContext(context.Background(), "tcp", f.listener.Addr().String()) + if err != nil { + t.Fatal(err) + } + f.peers = append(f.peers, peer) + select { + case conn := <-f.listener.accepted: + return peer, conn + case <-time.After(benchmarkCancelJoinBudget): + t.Fatal("owned TCP socket was not accepted") + return nil, nil + } +} + +func benchmarkCancelJoin(t *testing.T, done <-chan struct{}, name string) { + t.Helper() + select { + case <-done: + case <-time.After(benchmarkCancelJoinBudget): + t.Errorf("%s did not join", name) + } +} + +func benchmarkCancelWaitPhase(t *testing.T, phase <-chan struct{}, name string) { + t.Helper() + select { + case <-phase: + case <-time.After(benchmarkCancelJoinBudget): + t.Fatalf("%s was not reached", name) + } +} + +func benchmarkCancelRequirePeerClosed(t *testing.T, peer net.Conn, name string) { + t.Helper() + if err := peer.SetReadDeadline(time.Now().Add(200 * time.Millisecond)); err != nil { + t.Fatal(err) + } + var one [1]byte + n, err := peer.Read(one[:]) + if n != 0 || (!errors.Is(err, io.EOF) && !errors.Is(err, syscall.ECONNRESET)) { + t.Errorf("%s remains open before independent cleanup: n=%d err=%v", name, n, err) + } +} + +func benchmarkCancelRequireServerReturned(t *testing.T, f *benchmarkCancelServerFixture) { + t.Helper() + select { + case <-f.done: + if err := <-f.result; err != nil { + t.Errorf("Serve result = %v, want nil for cancellation/closed listener", err) + } + case <-time.After(benchmarkCancelReturnBudget): + t.Error("Serve still waits for an accepted worker after shutdown") + } +} + +func TestBenchmarkServerCancellationClosesAcceptedSockets(t *testing.T) { + for _, phase := range []string{"header", "upload", "pending_admission"} { + t.Run(phase, func(t *testing.T) { + f := benchmarkCancelStartServer(t) + peer, conn := f.connect(t) + benchmarkCancelWaitPhase(t, conn.readStarted, "actual header Read") + if phase == "upload" { + if err := peer.SetWriteDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + if _, err := peer.Write(benchmarkCancelHeader(ModeUpload, 1024)); err != nil { + t.Fatal(err) + } + benchmarkCancelWaitPhase(t, conn.uploadStarted, "actual upload payload Read") + } + var pending net.Conn + if phase == "pending_admission" { + var pendingConn *benchmarkCancelServerConn + pending, pendingConn = f.connect(t) + select { + case <-pendingConn.readStarted: + t.Fatal("second socket reached handler despite full admission") + default: + } + } + f.cancel() + benchmarkCancelRequirePeerClosed(t, peer, "active TCP socket") + if pending != nil { + benchmarkCancelRequirePeerClosed(t, pending, "pending admission TCP socket") + } + benchmarkCancelRequireServerReturned(t, f) + }) + } +} + +func TestBenchmarkServerListenerCloseJoinsActiveWorker(t *testing.T) { + f := benchmarkCancelStartServer(t) + peer, conn := f.connect(t) + benchmarkCancelWaitPhase(t, conn.readStarted, "actual header Read") + if err := f.listener.Close(); err != nil { + t.Fatal(err) + } + // The parent context is deliberately live here. Returning from Accept + // must also stop owned workers and the listener cancellation watcher. + benchmarkCancelRequirePeerClosed(t, peer, "externally stopped server TCP socket") + benchmarkCancelRequireServerReturned(t, f) +} + +func benchmarkCancelHeader(mode byte, size int64) []byte { + header := make([]byte, headerSize) + copy(header[:4], magic[:]) + header[4], header[5] = protocolVersion, mode + binary.BigEndian.PutUint64(header[8:], uint64(size)) + return header +} + +type benchmarkCancelRunConn struct { + net.Conn + headerWrite chan struct{} + payloadWrite chan struct{} + readStarted chan struct{} + writeCalls atomic.Int64 + headerOnce sync.Once + payloadOnce sync.Once + readOnce sync.Once + closeCalls atomic.Int64 +} + +func (c *benchmarkCancelRunConn) Write(p []byte) (int, error) { + if c.writeCalls.Add(1) == 1 { + c.headerOnce.Do(func() { close(c.headerWrite) }) + } else { + c.payloadOnce.Do(func() { close(c.payloadWrite) }) + } + return c.Conn.Write(p) +} + +func (c *benchmarkCancelRunConn) Read(p []byte) (int, error) { + c.readOnce.Do(func() { close(c.readStarted) }) + return c.Conn.Read(p) +} + +func (c *benchmarkCancelRunConn) Close() error { + c.closeCalls.Add(1) + return c.Conn.Close() +} + +func TestBenchmarkRunManualCancellation(t *testing.T) { + for _, deadline := range []string{"none", "distant"} { + t.Run(deadline, func(t *testing.T) { + for _, phase := range []string{"header", "download", "upload", "ack"} { + t.Run(phase, func(t *testing.T) { + benchmarkCancelRunAtPhase(t, deadline, phase) + }) + } + }) + } +} + +func benchmarkCancelRunAtPhase(t *testing.T, deadline, phase string) { + t.Helper() + base := context.Background() + var deadlineCancel context.CancelFunc + if deadline == "distant" { + base, deadlineCancel = context.WithTimeout(base, time.Minute) + defer deadlineCancel() + } + ctx, cancel := context.WithCancel(base) + defer cancel() + clientRaw, peer := net.Pipe() + conn := &benchmarkCancelRunConn{ + Conn: clientRaw, headerWrite: make(chan struct{}), payloadWrite: make(chan struct{}), readStarted: make(chan struct{}), + } + release := make(chan struct{}) + var releaseOnce sync.Once + peerDone := make(chan struct{}) + peerResult := make(chan error, 1) + runDone := make(chan struct{}) + type outcome struct { + result Result + err error + } + runResult := make(chan outcome, 1) + mode, size := ModeDownload, int64(1024) + if phase == "upload" { + mode = ModeUpload + } else if phase == "ack" { + size = 0 + } + t.Cleanup(func() { + cancel() + // Close both owned endpoints before any joins; do not rely on Run + // to honor cancellation or the deliberately distant deadline. + _ = clientRaw.Close() + _ = peer.Close() + releaseOnce.Do(func() { close(release) }) + benchmarkCancelJoin(t, runDone, "Run cleanup") + benchmarkCancelJoin(t, peerDone, "peer cleanup") + }) + go func() { + defer close(peerDone) + var err error + if phase != "header" { + header := make([]byte, headerSize) + _, err = io.ReadFull(peer, header) + if err == nil && string(header) != string(benchmarkCancelHeader(mode, size)) { + err = fmt.Errorf("wrong actual benchmark header: %x", header) + } + } + peerResult <- err + <-release + }() + dialer := transport.DialFunc(func(_ context.Context, network, address string) (net.Conn, error) { + if network != "tcp" || address != "owned-benchmark.invalid:1" { + return nil, fmt.Errorf("unexpected benchmark dial %q %q", network, address) + } + return conn, nil + }) + go func() { + defer close(runDone) + result, err := Run(ctx, dialer, "owned-benchmark.invalid:1", mode, size) + runResult <- outcome{result: result, err: err} + }() + var operation <-chan struct{} + switch phase { + case "header": + operation = conn.headerWrite + case "upload": + operation = conn.payloadWrite + default: + operation = conn.readStarted + } + benchmarkCancelWaitPhase(t, operation, "actual pending "+phase+" I/O") + // Set this peer-only oracle deadline before Run closes the other pipe + // endpoint: net.Pipe rejects SetReadDeadline after a remote Close. + if err := peer.SetReadDeadline(time.Now().Add(benchmarkCancelReturnBudget + 200*time.Millisecond)); err != nil { + t.Fatal(err) + } + // No peer operation can complete the selected I/O. Manual cancellation + // must stop Run before an endpoint is closed by independent cleanup. + cancel() + select { + case <-runDone: + got := <-runResult + if !errors.Is(got.err, ctx.Err()) { + t.Errorf("manual cancel at %s = %v, want %v identity", phase, got.err, ctx.Err()) + } + if got.result != (Result{}) { + t.Errorf("canceled transfer reported success: %+v", got.result) + } + if conn.closeCalls.Load() == 0 { + t.Error("Run returned without closing its owned connection") + } + var one [1]byte + if n, err := peer.Read(one[:]); n != 0 || !errors.Is(err, io.EOF) { + t.Errorf("Run peer not closed before cleanup: n=%d err=%v", n, err) + } + case <-time.After(benchmarkCancelReturnBudget): + t.Errorf("manual cancel at %s (deadline=%s) leaves Run blocked", phase, deadline) + } + select { + case err := <-peerResult: + if err != nil { + t.Errorf("controlled peer header = %v", err) + } + case <-time.After(benchmarkCancelJoinBudget): + t.Error("controlled peer did not observe the complete header") + } +} diff --git a/internal/netbench/dial_ownership_test.go b/internal/netbench/dial_ownership_test.go new file mode 100644 index 0000000..5b6c1fb --- /dev/null +++ b/internal/netbench/dial_ownership_test.go @@ -0,0 +1,210 @@ +package netbench + +import ( + "context" + "errors" + "io" + "net" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +type benchmarkCloseCountConn struct { + net.Conn + closes atomic.Int32 +} + +func (c *benchmarkCloseCountConn) Close() error { + c.closes.Add(1) + return c.Conn.Close() +} + +func TestBenchmarkRunOwnsFailedDialResult(t *testing.T) { + for _, canceled := range []bool{false, true} { + name := "live" + if canceled { + name = "canceled during dial" + } + t.Run(name, func(t *testing.T) { + client, peer := net.Pipe() + defer peer.Close() + conn := &benchmarkCloseCountConn{Conn: client} + defer client.Close() // Independent cleanup if Run loses ownership. + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + cause := errors.New("selected dial failure") + dialer := transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + if canceled { + cancel() + } + return conn, cause + }) + result, err := Run(ctx, dialer, "unused", ModeDownload, 1) + if err != cause || result != (Result{}) { + t.Errorf("Run = %+v, %v; want zero result and original dial error", result, err) + } + if got := conn.closes.Load(); got != 1 { + t.Errorf("failed dial result closed %d times, want 1", got) + } + }) + } +} + +func TestBenchmarkRunRejectsNilDialResult(t *testing.T) { + dialer := transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + return nil, nil + }) + defer func() { + if recovered := recover(); recovered != nil { + t.Errorf("nil dial result panicked: %v", recovered) + } + }() + result, err := Run(context.Background(), dialer, "unused", ModeDownload, 1) + if err == nil || result != (Result{}) { + t.Errorf("Run = %+v, %v; want zero result and an error", result, err) + } +} + +func TestBenchmarkRunSkipsCanceledDial(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + var calls atomic.Int32 + dialer := transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + calls.Add(1) + return nil, errors.New("dial should not be called") + }) + result, err := Run(ctx, dialer, "unused", ModeDownload, 1) + if !errors.Is(err, context.Canceled) || result != (Result{}) { + t.Errorf("Run = %+v, %v; want zero result and context cancellation", result, err) + } + if got := calls.Load(); got != 0 { + t.Errorf("already canceled context made %d dial calls", got) + } +} + +func TestBenchmarkRunOwnsLateSuccessfulDial(t *testing.T) { + client, peer := net.Pipe() + defer client.Close() + defer peer.Close() + conn := &benchmarkCloseCountConn{Conn: client} + ctx, cancel := context.WithCancelCause(context.Background()) + defer cancel(nil) + cause := errors.New("selected cancellation cause") + dialer := transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + cancel(cause) + return conn, nil + }) + result, err := Run(ctx, dialer, "unused", ModeDownload, 1) + if err != cause || result != (Result{}) { + t.Errorf("Run = %+v, %v; want zero result and original cancellation cause", result, err) + } + if got := conn.closes.Load(); got != 1 { + t.Errorf("late dial result closed %d times, want 1", got) + } +} + +type benchmarkJoinedCloseConn struct { + net.Conn + readStarted chan struct{} + readExited chan struct{} + closeStarted chan struct{} + closeFinished chan struct{} + releaseClose chan struct{} + closes atomic.Int32 +} + +func (c *benchmarkJoinedCloseConn) Read(p []byte) (int, error) { + close(c.readStarted) + n, err := c.Conn.Read(p) + close(c.readExited) + return n, err +} + +func (c *benchmarkJoinedCloseConn) Close() error { + c.closes.Add(1) + err := c.Conn.Close() + close(c.closeStarted) + <-c.releaseClose + close(c.closeFinished) + return err +} + +func TestBenchmarkRunJoinsCancellationClose(t *testing.T) { + client, peer := net.Pipe() + conn := &benchmarkJoinedCloseConn{ + Conn: client, readStarted: make(chan struct{}), readExited: make(chan struct{}), + closeStarted: make(chan struct{}), closeFinished: make(chan struct{}), releaseClose: make(chan struct{}), + } + ctx, cancel := context.WithCancelCause(context.Background()) + cause := errors.New("selected transfer cancellation cause") + var release sync.Once + peerDone, runDone := make(chan struct{}), make(chan struct{}) + peerResult, runResult := make(chan error, 1), make(chan error, 1) + t.Cleanup(func() { + cancel(nil) + _ = client.Close() + _ = peer.Close() + release.Do(func() { close(conn.releaseClose) }) + benchmarkCancelJoin(t, runDone, "joined-close Run cleanup") + benchmarkCancelJoin(t, peerDone, "joined-close peer cleanup") + }) + go func() { + defer close(peerDone) + header := make([]byte, headerSize) + _, err := io.ReadFull(peer, header) + peerResult <- err + }() + dialer := transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { return conn, nil }) + go func() { + defer close(runDone) + result, err := Run(ctx, dialer, "unused", ModeDownload, 1) + if result != (Result{}) { + t.Errorf("canceled transfer reported a result: %+v", result) + } + runResult <- err + }() + benchmarkCancelWaitPhase(t, conn.readStarted, "actual download Read") + benchmarkCancelJoin(t, peerDone, "complete header peer") + if err := <-peerResult; err != nil { + t.Fatal(err) + } + cancel(cause) + benchmarkCancelWaitPhase(t, conn.closeStarted, "cancellation Close entry") + benchmarkCancelWaitPhase(t, conn.readExited, "interrupted download Read return") + select { + case <-runDone: + t.Error("Run returned before its cancellation Close completed") + default: + } + release.Do(func() { close(conn.releaseClose) }) + benchmarkCancelWaitPhase(t, runDone, "Run after Close completion") + if err := <-runResult; err != cause { + t.Errorf("Run cancellation = %v, want original cause", err) + } + select { + case <-conn.closeFinished: + default: + t.Error("Run returned without joining Close") + } + if got := conn.closes.Load(); got != 1 { + t.Errorf("cancellation closed the connection %d times, want 1", got) + } +} + +func TestBenchmarkRunSkipsExpiredDeadline(t *testing.T) { + ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second)) + defer cancel() + var calls atomic.Int32 + dialer := transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + calls.Add(1) + return nil, errors.New("dial should not be called") + }) + result, err := Run(ctx, dialer, "unused", ModeDownload, 1) + if !errors.Is(err, context.DeadlineExceeded) || result != (Result{}) || calls.Load() != 0 { + t.Errorf("expired deadline Run = %+v, %v, %d dial calls", result, err, calls.Load()) + } +} diff --git a/internal/netbench/listener_ownership_test.go b/internal/netbench/listener_ownership_test.go new file mode 100644 index 0000000..f5b3421 --- /dev/null +++ b/internal/netbench/listener_ownership_test.go @@ -0,0 +1,128 @@ +package netbench + +import ( + "context" + "errors" + "net" + "sync" + "sync/atomic" + "testing" + "time" +) + +// Only the second Accept error is synthetic. The first accepted socket and +// every Close below are real loopback TCP operations with their original result. +type benchmarkErrorListener struct { + net.Listener + failure error + accepted chan *benchmarkCancelServerConn + secondAccept chan struct{} + releaseAccept chan struct{} + acceptCalls atomic.Int32 + closeCalls atomic.Int32 + mu sync.Mutex + closed bool + conn net.Conn +} + +func (l *benchmarkErrorListener) Accept() (net.Conn, error) { + if l.acceptCalls.Add(1) != 1 { + close(l.secondAccept) + <-l.releaseAccept + return nil, l.failure + } + raw, err := l.Listener.Accept() + if err != nil { + return nil, err + } + conn := &benchmarkCancelServerConn{ + Conn: raw, readStarted: make(chan struct{}), uploadStarted: make(chan struct{}), + } + l.mu.Lock() + if l.closed { + l.mu.Unlock() + _ = raw.Close() + return nil, net.ErrClosed + } + l.conn = conn + l.mu.Unlock() + l.accepted <- conn // Exactly one accepted connection; channel capacity is one. + return conn, nil +} + +func (l *benchmarkErrorListener) Close() error { + l.closeCalls.Add(1) + return l.Listener.Close() +} + +func (l *benchmarkErrorListener) forceCloseConn() { + l.mu.Lock() + l.closed = true + conn := l.conn + l.mu.Unlock() + if conn != nil { + _ = conn.Close() + } +} + +func TestBenchmarkServerAcceptErrorOwnsListenerAndWorker(t *testing.T) { + raw, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + failure := errors.New("controlled second Accept failure") + ln := &benchmarkErrorListener{ + Listener: raw, failure: failure, accepted: make(chan *benchmarkCancelServerConn, 1), + secondAccept: make(chan struct{}), releaseAccept: make(chan struct{}), + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan struct{}) + result := make(chan error, 1) + var releaseOnce sync.Once + var peer net.Conn + t.Cleanup(func() { + // Bypass the observed wrapper for independent listener cleanup; its + // Close counter measures production ownership only. + cancel() + _ = raw.Close() + ln.forceCloseConn() + if peer != nil { + _ = peer.Close() + } + releaseOnce.Do(func() { close(ln.releaseAccept) }) + benchmarkCancelJoin(t, done, "Accept-error Serve cleanup") + }) + go func() { + defer close(done) + result <- (&Server{MaxConnections: 1}).Serve(ctx, ln) + }() + peer, err = (&net.Dialer{Timeout: time.Second}).DialContext(context.Background(), "tcp", raw.Addr().String()) + if err != nil { + t.Fatal(err) + } + var conn *benchmarkCancelServerConn + select { + case conn = <-ln.accepted: + case <-time.After(benchmarkCancelJoinBudget): + t.Fatal("real TCP socket was not accepted") + } + benchmarkCancelWaitPhase(t, conn.readStarted, "active real header Read") + benchmarkCancelWaitPhase(t, ln.secondAccept, "controlled second Accept") + releaseOnce.Do(func() { close(ln.releaseAccept) }) + benchmarkCancelRequirePeerClosed(t, peer, "Accept-error worker TCP socket") + select { + case <-done: + if got := <-result; got != failure { + t.Errorf("Serve changed original Accept error identity: got %v, want %v", got, failure) + } + case <-time.After(benchmarkCancelReturnBudget): + t.Error("Serve did not join its active worker after Accept failure") + } + if got := ln.closeCalls.Load(); got != 1 { + t.Errorf("production listener Close calls = %d, want exactly one before cleanup", got) + } + if ctx.Err() != nil { + t.Errorf("Serve canceled its caller-owned parent context: %v", ctx.Err()) + } +} diff --git a/internal/netbench/netbench.go b/internal/netbench/netbench.go index cce534c..8017eab 100644 --- a/internal/netbench/netbench.go +++ b/internal/netbench/netbench.go @@ -33,7 +33,8 @@ type Server struct { Timeout time.Duration } -// Serve accepts connections until ctx is canceled or the listener fails. +// Serve owns the listener and accepted connections until ctx is canceled or +// the listener fails. It closes and joins active transfers before returning. func (s *Server) Serve(ctx context.Context, ln net.Listener) error { maxConnections := s.MaxConnections if maxConnections <= 0 { @@ -41,13 +42,25 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error { } sem := make(chan struct{}, maxConnections) var wg sync.WaitGroup - - go func() { - <-ctx.Done() + serveCtx, cancel := context.WithCancel(ctx) + listenerClosed := make(chan struct{}) + stopListener := context.AfterFunc(serveCtx, func() { _ = ln.Close() - }() + close(listenerClosed) + }) - defer wg.Wait() + defer func() { + // Accept errors also end this ownership scope, even with a live + // parent. Cancel workers before joining them, not after their I/O + // eventually reaches the independent transfer deadline. + cancel() + if stopListener() { + _ = ln.Close() + } else { + <-listenerClosed + } + wg.Wait() + }() for { conn, err := ln.Accept() if err != nil { @@ -58,7 +71,7 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error { } select { case sem <- struct{}{}: - case <-ctx.Done(): + case <-serveCtx.Done(): _ = conn.Close() return nil } @@ -66,7 +79,8 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error { go func() { defer wg.Done() defer func() { <-sem }() - defer conn.Close() + finish := ownBenchmarkConnection(serveCtx, conn) + defer finish() _ = s.handle(conn) }() } @@ -127,7 +141,9 @@ func (r Result) Mbps() float64 { return float64(r.Bytes*8) / r.Duration.Seconds() / 1_000_000 } -// Run performs one upload or download through dialer. +// Run performs one upload or download through dialer. Cancellation interrupts +// the owned transfer, including its header and completion acknowledgement; +// it does not close the dialer or its shared physical tunnel connection. func Run(ctx context.Context, dialer transport.Dialer, target string, mode byte, size int64) (Result, error) { if mode != ModeDownload && mode != ModeUpload { return Result{}, errors.New("invalid benchmark mode") @@ -135,14 +151,33 @@ func Run(ctx context.Context, dialer transport.Dialer, target string, mode byte, if size < 0 { return Result{}, errors.New("size must not be negative") } + if err := benchmarkContextError(ctx); err != nil { + return Result{}, err + } conn, err := dialer.DialContext(ctx, "tcp", target) - if err != nil { + if err != nil || conn == nil { + if conn != nil { + _ = conn.Close() + } + if err == nil { + err = errors.New("benchmark dialer returned a nil connection") + } + return Result{}, err + } + finish := ownBenchmarkConnection(ctx, conn) + defer finish() + if err := benchmarkContextError(ctx); err != nil { return Result{}, err } - defer conn.Close() if deadline, ok := ctx.Deadline(); ok { _ = conn.SetDeadline(deadline) } + ioError := func(err error) (Result, error) { + if canceled := benchmarkContextError(ctx); canceled != nil { + err = canceled + } + return Result{}, err + } header := make([]byte, headerSize) copy(header[:4], magic[:]) @@ -150,22 +185,25 @@ func Run(ctx context.Context, dialer transport.Dialer, target string, mode byte, header[5] = mode binary.BigEndian.PutUint64(header[8:], uint64(size)) if _, err := conn.Write(header); err != nil { - return Result{}, err + return ioError(err) } start := time.Now() switch mode { case ModeDownload: if _, err := io.CopyN(io.Discard, conn, size); err != nil { - return Result{}, err + return ioError(err) } case ModeUpload: if err := writeZeros(conn, size); err != nil { - return Result{}, err + return ioError(err) } } var ack [1]byte if _, err := io.ReadFull(conn, ack[:]); err != nil { + return ioError(err) + } + if err := benchmarkContextError(ctx); err != nil { return Result{}, err } if ack[0] != 0 { @@ -174,6 +212,36 @@ func Run(ctx context.Context, dialer transport.Dialer, target string, mode byte, return Result{Mode: mode, Bytes: size, Duration: time.Since(start)}, nil } +// Close exactly this owned connection once. A cancellation callback can still +// be executing when stop returns false; join it before releasing the worker or +// returning a result. Normal completion stops the callback and closes directly. +func ownBenchmarkConnection(ctx context.Context, conn net.Conn) func() { + closed := make(chan struct{}) + stop := context.AfterFunc(ctx, func() { + _ = conn.Close() + close(closed) + }) + return func() { + if stop() { + _ = conn.Close() + } else { + <-closed + } + } +} + +func benchmarkContextError(ctx context.Context) error { + if err := context.Cause(ctx); err != nil { + return err + } + // The socket deadline may fire just before the context timer publishes + // Done. Preserve the declared context deadline in that narrow race too. + if deadline, ok := ctx.Deadline(); ok && !time.Now().Before(deadline) { + return context.DeadlineExceeded + } + return nil +} + func writeZeros(w io.Writer, size int64) error { buf := make([]byte, 64<<10) for size > 0 { diff --git a/internal/tunnel/tunnel_test.go b/internal/tunnel/tunnel_test.go index a7880fe..c317e31 100644 --- a/internal/tunnel/tunnel_test.go +++ b/internal/tunnel/tunnel_test.go @@ -854,57 +854,189 @@ func TestHardenedQUICConfigDisablesReplayableFeaturesAndEnablesDatagrams(t *test } func TestQUICConnectionLimitAndPreAuthenticationTimeout(t *testing.T) { - serverTLS, clientTLS := testTLSConfigs(t) - server, err := ListenQUIC(QUICServerConfig{ - Address: "127.0.0.1:0", - Token: testToken, - TLSConfig: serverTLS, - HandshakeTimeout: 500 * time.Millisecond, - MaxConnections: 1, - }) - if err != nil { - t.Fatal(err) + newServer := func(t *testing.T, handshakeTimeout time.Duration) (*QUICServer, *tls.Config, func()) { + t.Helper() + serverTLS, clientTLS := testTLSConfigs(t) + server, err := ListenQUIC(QUICServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + HandshakeTimeout: handshakeTimeout, MaxConnections: 1, + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + serveDone := make(chan struct{}) + serveResult := make(chan error, 1) + started := false + t.Cleanup(func() { + cancel() + closeDone := make(chan struct{}) + closeResult := make(chan error, 1) + go func() { + defer close(closeDone) + closeResult <- server.Close() + }() + select { + case <-closeDone: + if err := <-closeResult; err != nil { + t.Errorf("Close: %v", err) + } + case <-time.After(2 * time.Second): + t.Error("server Close did not join") + } + if started { + select { + case <-serveDone: + if err := <-serveResult; err != nil { + t.Errorf("Serve: %v", err) + } + case <-time.After(2 * time.Second): + t.Error("server Serve did not join") + } + } + }) + rawTLS, err := clientTLSConfig(clientTLS, server.Addr().String()) + if err != nil { + t.Fatal(err) + } + start := func() { + if started { + t.Fatal("fixture Serve started twice") + } + started = true + go func() { + defer close(serveDone) + serveResult <- server.Serve(ctx) + }() + } + return server, rawTLS, start } - ctx, cancel := context.WithCancel(context.Background()) - serveDone := make(chan error, 1) - go func() { serveDone <- server.Serve(ctx) }() - t.Cleanup(func() { - cancel() - _ = server.Close() - if err := <-serveDone; err != nil { - t.Errorf("Serve: %v", err) + dial := func(t *testing.T, server *QUICServer, config *tls.Config) (*quic.Conn, error) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + conn, err := quic.DialAddr(ctx, server.Addr().String(), config.Clone(), hardenedQUICClientConfig(nil)) + if conn != nil { + t.Cleanup(func() { + _ = conn.CloseWithError(0, "test complete") + select { + case <-conn.Context().Done(): + case <-time.After(2 * time.Second): + t.Error("owned client connection did not close") + } + }) } - }) - - rawClientTLS, err := clientTLSConfig(clientTLS, server.Addr().String()) - if err != nil { - t.Fatal(err) + return conn, err } - first, err := quic.DialAddr(context.Background(), server.Addr().String(), rawClientTLS.Clone(), hardenedQUICClientConfig(nil)) - if err != nil { - t.Fatalf("first QUIC connection: %v", err) + assertRemoteError := func(t *testing.T, err error, code quic.ApplicationErrorCode, message string) { + t.Helper() + var applicationErr *quic.ApplicationError + if !errors.As(err, &applicationErr) || !applicationErr.Remote || applicationErr.ErrorCode != code || applicationErr.ErrorMessage != message { + t.Fatalf("remote QUIC rejection = %v, want code %#x and message %q", err, code, message) + } } - defer first.CloseWithError(0, "test complete") - second, secondErr := quic.DialAddr(context.Background(), server.Addr().String(), rawClientTLS.Clone(), hardenedQUICClientConfig(nil)) - if secondErr == nil { - defer second.CloseWithError(0, "test complete") + t.Run("global_limit", func(t *testing.T) { + // This fixture tests admission, not its separate 500ms auth timer. + // Use the normal default so the acceptance barrier and rejection window + // cannot accidentally compete with a short pre-authentication expiry. + server, rawTLS, start := newServer(t, 0) + first, err := dial(t, server, rawTLS) + if err != nil { + t.Fatalf("first QUIC connection: %v", err) + } + // quic-go runs the TLS handshake independently of application Serve. + // Actually demonstrate that a completed client handshake is not a + // server-admission witness before establishing the required barrier. + server.connMu.Lock() + unregistered := len(server.conns) == 0 + server.connMu.Unlock() + if !unregistered || len(server.connSem) != 0 { + t.Fatal("fixture admitted a connection before starting Serve") + } + start() + local, ok := first.LocalAddr().(*net.UDPAddr) + if !ok { + t.Fatalf("first local address type = %T", first.LocalAddr()) + } + barrierTimeout := time.NewTimer(time.Second) + defer barrierTimeout.Stop() + tick := time.NewTicker(time.Millisecond) + defer tick.Stop() + for { + server.connMu.Lock() + matched := false + if len(server.conns) == 1 { + for conn := range server.conns { + remote, ok := conn.RemoteAddr().(*net.UDPAddr) + // DialAddr binds a wildcard UDP socket; on this exclusively + // loopback fixture, its source port identifies the exact peer. + matched = ok && remote.IP.IsLoopback() && remote.Port == local.Port + } + } + admitted := matched && len(server.connSem) == 1 + server.connMu.Unlock() + if admitted { + break + } + select { + case <-first.Context().Done(): + t.Fatalf("first connection closed before admission: %v", context.Cause(first.Context())) + case <-barrierTimeout.C: + t.Fatal("first connection did not own the server's single admission slot") + case <-tick.C: + } + } + second, rejection := dial(t, server, rawTLS) + if rejection == nil { + select { + case <-second.Context().Done(): + rejection = context.Cause(second.Context()) + case <-time.After(250 * time.Millisecond): + t.Fatal("connection above MaxConnections was not rejected") + } + } + // The same source also has a limit; its identically coded rejection + // must not substitute for proof that the global admission branch ran. + assertRemoteError(t, rejection, connectionRejected, "connection limit reached") + select { + case <-first.Context().Done(): + t.Fatalf("global admission rejected the already-owned first connection: %v", context.Cause(first.Context())) + default: + } + }) + + t.Run("pre_authentication_timeout", func(t *testing.T) { + const preAuthBudget = 500 * time.Millisecond + server, rawTLS, start := newServer(t, preAuthBudget) + // Complete TLS before the application can start its auth timer. This + // independently preserves the healthy handshake / no-premature-close + // control without guessing when a concurrently started worker ran. + conn, err := dial(t, server, rawTLS) + if err != nil { + t.Fatalf("unauthenticated QUIC handshake: %v", err) + } select { - case <-second.Context().Done(): - case <-time.After(250 * time.Millisecond): - t.Fatal("connection above MaxConnections was not rejected") + case <-conn.Context().Done(): + t.Fatalf("unauthenticated connection closed before Serve: %v", context.Cause(conn.Context())) + default: } - } - select { - case <-first.Context().Done(): - t.Fatal("first unauthenticated connection closed before its authentication deadline") - default: - } - select { - case <-first.Context().Done(): - case <-time.After(2 * time.Second): - t.Fatal("unauthenticated QUIC connection survived its authentication deadline") - } + began := time.Now() + start() + select { + case <-conn.Context().Done(): + case <-time.After(2 * time.Second): + t.Fatal("unauthenticated QUIC connection survived its authentication deadline") + } + elapsed := time.Since(began) + assertRemoteError(t, context.Cause(conn.Context()), authenticationTimeout, "authentication timeout") + // The server-side timer cannot start before this Serve invocation. + // Therefore a shorter elapsed interval is an actual premature expiry, + // not merely a delayed client observation or admission-barrier race. + if elapsed < preAuthBudget { + t.Fatalf("authentication timeout arrived after %s, before the %s budget", elapsed, preAuthBudget) + } + }) } func TestAuthenticatedQUICConnectionSurvivesPreAuthenticationTimeout(t *testing.T) {