From 4c47c00b0dc0d692d63c09e309f190b2ba629278 Mon Sep 17 00:00:00 2001 From: cppla Date: Sat, 3 Oct 2026 08:04:19 +0800 Subject: [PATCH 1/2] fix(web): share H2 initialization and join client closure --- docs/WEB_COVER.md | 20 +- internal/tunnel/web_h2_client.go | 233 ++++-- internal/tunnel/web_h2_close_join_test.go | 437 +++++++++++ internal/tunnel/web_h2_dial_result_test.go | 114 +++ internal/tunnel/web_h2_dial_sharing_test.go | 706 ++++++++++++++++++ internal/tunnel/web_h2_initialization_test.go | 187 ++++- 6 files changed, 1613 insertions(+), 84 deletions(-) create mode 100644 internal/tunnel/web_h2_close_join_test.go create mode 100644 internal/tunnel/web_h2_dial_result_test.go create mode 100644 internal/tunnel/web_h2_dial_sharing_test.go diff --git a/docs/WEB_COVER.md b/docs/WEB_COVER.md index fa8ab58..05b26fd 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -380,15 +380,27 @@ IPv4 candidates and starts them with a short stagger; each candidate retains its own UDP socket so a blackholed first address cannot consume the entire H3 budget before a working family is tried. -Source builds after v1.0.1 share the result of a cold H3 physical connection +Source builds after v1.0.1 share the result of a cold H2 or H3 physical connection attempt with all callers already waiting on it. A failed handshake does not make those callers start replacement handshakes one after another. A later invocation may retry, subject to the existing `web-auto` cooldown policy. The shared attempt belongs to the client and retains its configured physical dial timeout; canceling one caller stops only that caller's wait, not the -attempt needed by other callers. If all callers abandon it, the attempt can -continue until its existing timeout or client shutdown. Client Close cancels -and joins owned dialing and session-cleanup workers before returning. This +attempt needed by other callers. H2 keeps its TCP-connect timeout separate +from the existing TLS/H2 initialization budget; its bounded TLS retry after +HelloRetryRequest still uses the original remaining initialization budget. +If all callers abandon it, initialization can continue until its configured +timeout or client shutdown. A successful unclaimed H2 connection can remain +pooled, but does not send CONNECT or generate application authentication until +a live caller reserves its first stream. That caller alone owns bootstrap; +subsequent callers retain the existing short connection-bound credentials. +Client Close cancels and joins owned dialing and detached session-cleanup +workers before returning. Concurrent H2 Close callers wait for the same +completed closure and result, including late physical dial results. This +does not claim synchronous termination of every dependency goroutine; custom +TLS callbacks must return, and operations ignoring cancellation cannot be +forcibly interrupted. A failed TCP dial's non-nil connection is closed before +its failure is published; a nil successful result fails closed. This does not change authentication, wire profiles, timeout defaults or sibling stream ownership, and is not a general connection-speed or browser-equivalence claim. The published v1.0.1 binary does not contain this change. diff --git a/internal/tunnel/web_h2_client.go b/internal/tunnel/web_h2_client.go index dcb2378..c17c13e 100644 --- a/internal/tunnel/web_h2_client.go +++ b/internal/tunnel/web_h2_client.go @@ -30,7 +30,10 @@ type WebH2ClientConfig struct { // reference implemented by uTLS v1.8.2. Native is for tests/debugging. FingerprintProfile FingerprintProfile HandshakeTimeout time.Duration - DialTimeout time.Duration + // DialTimeout bounds each TCP connect, separately from TLS/H2 setup. + // A cold physical attempt is client-owned; caller cancellation only stops + // that caller's wait. Zero retains the default TCP dial budget. + DialTimeout time.Duration // WriteByteTimeout limits stalled writes on the shared HTTP/2 connection. // Zero uses a conservative thirty-second default; negative values are // invalid. This is a TLS write-call budget, not a precise TCP byte-idle @@ -47,6 +50,7 @@ type WebH2Client struct { address string tlsConfig *tls.Config fingerprint FingerprintProfile + dialTimeout time.Duration handshakeTimeout time.Duration dialer transport.Dialer auth *webAuthSigner @@ -64,6 +68,19 @@ type WebH2Client struct { selected bool current *webH2ClientSession sessions map[*webH2ClientSession]struct{} + dial *webH2DialAttempt + // Additions are made under mu before the closed gate. Close joins both + // unpublished physical setup and cleanup detached from sessions. + workers sync.WaitGroup + closeDone chan struct{} + closeErr error +} + +// One immutable physical result is shared by callers already waiting for it. +// A caller owns its wait and CONNECT, never this client-owned initialization. +type webH2DialAttempt struct { + done chan struct{} + err error } type webH2ClientSession struct { @@ -78,13 +95,16 @@ type webH2ClientSession struct { opening int // active counts returned net.Conns until full close, including deadline // cancellation. Half-closes keep the opposite direction's ownership. - active int + active int + closeOnce sync.Once + closeErr error } type webH2ClientAuthState uint8 const ( - webH2ClientAuthBootstrapping webH2ClientAuthState = iota + webH2ClientAuthFresh webH2ClientAuthState = iota + webH2ClientAuthBootstrapping webH2ClientAuthReady webH2ClientAuthFailed ) @@ -153,6 +173,7 @@ func newWebH2ClientWithSigner(config WebH2ClientConfig, auth *webAuthSigner, cla address: config.ServerAddress, tlsConfig: tlsConfig, fingerprint: fingerprint, + dialTimeout: dialTimeout, handshakeTimeout: handshakeTimeout, dialer: &net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}, auth: auth, @@ -168,6 +189,7 @@ func newWebH2ClientWithSigner(config WebH2ClientConfig, auth *webAuthSigner, cla cancel: cancel, dialGate: make(chan struct{}, 1), sessions: make(map[*webH2ClientSession]struct{}), + closeDone: make(chan struct{}), }, nil } @@ -187,7 +209,7 @@ func (c *WebH2Client) DialContext(ctx context.Context, network, address string) if ctx == nil { return nil, errors.New("tunnel: nil dial context") } - if err := ctx.Err(); err != nil { + if err := contextError(ctx); err != nil { return nil, err } c.mu.Lock() @@ -329,10 +351,11 @@ func (c *WebH2Client) DialContext(ctx context.Context, network, address string) fail() return nil, context.DeadlineExceeded } - if !stopDial() && ctx.Err() != nil { + stopDial() + if err := contextError(ctx); err != nil { _ = response.Body.Close() fail() - return nil, ctx.Err() + return nil, err } c.mu.Lock() @@ -386,6 +409,9 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati select { case c.dialGate <- struct{}{}: case <-ctx.Done(): + if c.ctx.Err() != nil { + return webH2SessionReservation{}, net.ErrClosed + } return webH2SessionReservation{}, context.Cause(ctx) case <-c.ctx.Done(): return webH2SessionReservation{}, net.ErrClosed @@ -394,9 +420,9 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati <-c.dialGate return webH2SessionReservation{}, net.ErrClosed } - if err := ctx.Err(); err != nil { + if err := contextError(ctx); err != nil { <-c.dialGate - return webH2SessionReservation{}, context.Cause(ctx) + return webH2SessionReservation{}, err } c.mu.Lock() @@ -405,8 +431,24 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati <-c.dialGate return webH2SessionReservation{}, net.ErrClosed } + if err := contextError(ctx); err != nil { + c.mu.Unlock() + <-c.dialGate + return webH2SessionReservation{}, err + } if session := c.current; session != nil { switch session.authState { + case webH2ClientAuthFresh: + if session.h2.CanTakeNewRequest() { + // Claim bootstrap only when a live caller reserves a stream, + // not when a background worker publishes a cold connection. + session.authState = webH2ClientAuthBootstrapping + session.opening++ + c.mu.Unlock() + <-c.dialGate + return webH2SessionReservation{session: session, bootstrap: true}, nil + } + c.current = nil case webH2ClientAuthReady: if session.h2.CanTakeNewRequest() { session.opening++ @@ -424,6 +466,9 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati case <-ready: continue case <-ctx.Done(): + if c.ctx.Err() != nil { + return webH2SessionReservation{}, net.ErrClosed + } return webH2SessionReservation{}, context.Cause(ctx) case <-c.ctx.Done(): return webH2SessionReservation{}, net.ErrClosed @@ -433,34 +478,64 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati } } retired := c.cleanupIdleSessionsLocked() + attempt := c.dial + if attempt == nil { + attempt = &webH2DialAttempt{done: make(chan struct{})} + c.dial = attempt + c.workers.Add(1) + go c.runSessionDial(attempt) + } c.mu.Unlock() - closeWebH2Sessions(retired) - - session, err := c.openSession(ctx) - if err != nil { - <-c.dialGate + <-c.dialGate + c.closeRetiredSessions(retired) + select { + case <-attempt.done: + case <-c.ctx.Done(): + return webH2SessionReservation{}, net.ErrClosed + case <-ctx.Done(): + if c.ctx.Err() != nil { + return webH2SessionReservation{}, net.ErrClosed + } + return webH2SessionReservation{}, context.Cause(ctx) + } + if c.ctx.Err() != nil { + return webH2SessionReservation{}, net.ErrClosed + } + if err := contextError(ctx); err != nil { return webH2SessionReservation{}, err } - if !session.h2.CanTakeNewRequest() { - <-c.dialGate - _ = closeWebH2Session(session) - return webH2SessionReservation{}, errors.New("tunnel: new web-cover HTTP/2 connection rejected its first stream") + if attempt.err != nil { + return webH2SessionReservation{}, attempt.err } + } +} - c.mu.Lock() - if c.closed { - c.mu.Unlock() - <-c.dialGate - _ = closeWebH2Session(session) - return webH2SessionReservation{}, net.ErrClosed - } +func (c *WebH2Client) runSessionDial(attempt *webH2DialAttempt) { + defer c.workers.Done() + session, err := c.openSession(c.ctx) + if err == nil && !session.h2.CanTakeNewRequest() { + err = errors.New("tunnel: new web-cover HTTP/2 connection rejected its first stream") + } + c.mu.Lock() + accepted := !c.closed && err == nil + if c.closed { + err = net.ErrClosed + } + if accepted { c.current = session - session.opening++ c.sessions[session] = struct{}{} - c.mu.Unlock() - <-c.dialGate - return webH2SessionReservation{session: session, bootstrap: true}, nil } + c.mu.Unlock() + // Join late-result cleanup before either publishing failure or allowing + // Close to return. No session can register after the closed gate. + if session != nil && !accepted { + _ = closeWebH2Session(session) + } + c.mu.Lock() + attempt.err = err + c.dial = nil + close(attempt.done) + c.mu.Unlock() } func (c *WebH2Client) releaseSessionReservation(session *webH2ClientSession) { @@ -468,7 +543,7 @@ func (c *WebH2Client) releaseSessionReservation(session *webH2ClientSession) { session.opening-- retired := c.cleanupIdleSessionsLocked() c.mu.Unlock() - closeWebH2Sessions(retired) + c.closeRetiredSessions(retired) } func (c *WebH2Client) releaseSessionStream(session *webH2ClientSession) { @@ -476,7 +551,7 @@ func (c *WebH2Client) releaseSessionStream(session *webH2ClientSession) { session.active-- retired := c.cleanupIdleSessionsLocked() c.mu.Unlock() - closeWebH2Sessions(retired) + c.closeRetiredSessions(retired) } func (c *WebH2Client) openSession(ctx context.Context) (*webH2ClientSession, error) { @@ -487,14 +562,15 @@ func (c *WebH2Client) openSession(ctx context.Context) (*webH2ClientSession, err dialCancel() }() - raw, err := c.dialer.DialContext(dialCtx, "tcp", c.address) + raw, err := c.dialRaw(dialCtx) if err != nil { return nil, fmt.Errorf("tunnel: dial web-cover HTTP/2 server: %w", err) } // NewClientConn synchronously writes the HTTP/2 preface and SETTINGS after // TLS succeeds. Until that finishes, the connection is not in c.sessions, - // so Close cannot find it there. Keep raw I/O tied to both establishment - // contexts through initialization, not just through the TLS handshake. + // so pooled-session closure cannot find it there. The setup worker joins + // this watcher too; keep raw I/O tied to both establishment contexts + // through initialization, not just through the TLS handshake. initializationCtx, initializationCancel := context.WithTimeout(dialCtx, c.handshakeTimeout) defer initializationCancel() session, err := c.initializeSession(initializationCtx, raw, c.utlsSessionCache) @@ -508,7 +584,7 @@ func (c *WebH2Client) openSession(ctx context.Context) (*webH2ClientSession, err if cause := context.Cause(initializationCtx); cause != nil { return nil, fmt.Errorf("tunnel: web-cover HTTP/2 TLS handshake: %w", cause) } - raw, err = c.dialer.DialContext(initializationCtx, "tcp", c.address) + raw, err = c.dialRaw(initializationCtx) if err != nil { if cause := context.Cause(initializationCtx); cause != nil { err = cause @@ -518,6 +594,31 @@ func (c *WebH2Client) openSession(ctx context.Context) (*webH2ClientSession, err return c.initializeSession(initializationCtx, raw, nil) } +// The TCP budget stays separate from the existing TLS/H2 initialization +// budget. A retry after HRR is additionally bounded by that original budget. +func (c *WebH2Client) dialRaw(ctx context.Context) (net.Conn, error) { + timeout := c.dialTimeout + if timeout == 0 { + timeout = defaultDialTimeout + } + dialCtx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + raw, err := c.dialer.DialContext(dialCtx, "tcp", c.address) + if cause := contextError(dialCtx); cause != nil { + err = cause + } + if raw == nil && err == nil { + err = errors.New("tunnel: HTTP/2 physical dial returned no connection") + } + if err != nil { + if raw != nil { + err = errors.Join(err, raw.Close()) + } + return nil, err + } + return raw, nil +} + // initializeSession owns one physical attempt. Its immutable raw parameter // prevents a retired watcher's closure from affecting a replacement socket. func (c *WebH2Client) initializeSession(initializationCtx context.Context, raw net.Conn, sessionCache utls.ClientSessionCache) (*webH2ClientSession, error) { @@ -578,7 +679,7 @@ func (c *WebH2Client) initializeSession(initializationCtx context.Context, raw n raw: raw, conn: tlsConn, h2: clientConn, - authState: webH2ClientAuthBootstrapping, + authState: webH2ClientAuthFresh, authReady: make(chan struct{}), }, nil } @@ -617,7 +718,14 @@ func (c *WebH2Client) failSessionAuthentication(session *webH2ClientSession) { c.current = nil } delete(c.sessions, session) + tracked := !c.closed + if tracked { + c.workers.Add(1) + } c.mu.Unlock() + if tracked { + defer c.workers.Done() + } _ = closeWebH2Session(session) } @@ -628,7 +736,7 @@ func (c *WebH2Client) noteSessionFailure(session *webH2ClientSession) { } retired := c.cleanupIdleSessionsLocked() c.mu.Unlock() - closeWebH2Sessions(retired) + c.closeRetiredSessions(retired) } // Only inspect ownership maintained under c.mu. http2.ClientConn.State takes @@ -644,40 +752,57 @@ func (c *WebH2Client) cleanupIdleSessionsLocked() []*webH2ClientSession { delete(c.sessions, session) retired = append(retired, session) } + c.workers.Add(len(retired)) return retired } -func closeWebH2Sessions(sessions []*webH2ClientSession) { +func (c *WebH2Client) closeRetiredSessions(sessions []*webH2ClientSession) { for _, session := range sessions { _ = closeWebH2Session(session) + c.workers.Done() } } func closeWebH2Session(session *webH2ClientSession) error { - session.auth.close() - // Stop wire I/O first. Besides releasing HTTP/2 writers, this avoids a - // synchronous TLS close_notify waiting on an unresponsive peer. Retired - // sessions reach here only after all opening and returned streams drain. - var rawErr error - if session.raw != nil { - rawErr = session.raw.Close() - if errors.Is(rawErr, net.ErrClosed) { - rawErr = nil - } + if session == nil { + return nil } - return errors.Join(rawErr, session.h2.Close()) + session.closeOnce.Do(func() { + session.auth.close() + // Stop wire I/O first. Besides releasing HTTP/2 writers, this avoids a + // synchronous TLS close_notify waiting on an unresponsive peer. Retired + // sessions reach here only after all opening and returned streams drain. + var rawErr error + if session.raw != nil { + rawErr = session.raw.Close() + if errors.Is(rawErr, net.ErrClosed) { + rawErr = nil + } + } + var h2Err error + if session.h2 != nil { + h2Err = session.h2.Close() + } + session.closeErr = errors.Join(rawErr, h2Err) + }) + return session.closeErr } // Close prevents future dials, cancels active streams, and closes every pooled -// HTTP/2 connection. +// HTTP/2 connection. It joins owned setup and detached cleanup; concurrent +// callers observe the same completed closure and error. func (c *WebH2Client) Close() error { c.mu.Lock() + if c.closeDone == nil { + c.closeDone = make(chan struct{}) + } if c.closed { + done := c.closeDone c.mu.Unlock() - return nil + <-done + return c.closeErr } c.closed = true - c.cancel() sessions := make([]*webH2ClientSession, 0, len(c.sessions)) for session := range c.sessions { sessions = append(sessions, session) @@ -685,12 +810,16 @@ func (c *WebH2Client) Close() error { c.current = nil c.sessions = make(map[*webH2ClientSession]struct{}) c.mu.Unlock() + c.cancel() var result error for _, session := range sessions { result = errors.Join(result, closeWebH2Session(session)) } - return result + c.workers.Wait() + c.closeErr = result + close(c.closeDone) + return c.closeErr } var _ transport.Dialer = (*WebH2Client)(nil) diff --git a/internal/tunnel/web_h2_close_join_test.go b/internal/tunnel/web_h2_close_join_test.go new file mode 100644 index 0000000..6f2f1f1 --- /dev/null +++ b/internal/tunnel/web_h2_close_join_test.go @@ -0,0 +1,437 @@ +package tunnel + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "net/http" + "runtime" + "strings" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +// This is an explicit cleanup-completion barrier: the actual owned TCP socket +// closes first. It does not pretend that the socket's Close or I/O is stalled. +type webH2CloseJoinHeldConn struct { + net.Conn + entered chan struct{} + returned chan struct{} + release <-chan struct{} + once sync.Once + err error +} + +func (c *webH2CloseJoinHeldConn) Close() error { + c.once.Do(func() { + c.err = c.Conn.Close() + close(c.entered) + <-c.release + close(c.returned) + }) + return c.err +} + +func webH2CloseJoinCall(c *WebH2Client, started chan<- struct{}, results chan<- error, done chan struct{}) { + defer close(done) + started <- struct{}{} + results <- c.Close() +} + +func webH2CloseJoinCheckPending(t *testing.T, first, second <-chan struct{}) bool { + t.Helper() + timer := time.NewTimer(time.Second) + ticker := time.NewTicker(10 * time.Millisecond) + defer timer.Stop() + defer ticker.Stop() + for { + select { + case <-first: + t.Error("first Close returned before owned physical cleanup completed") + return false + case <-second: + t.Error("concurrent Close returned before owned physical cleanup completed") + return false + default: + } + stack := make([]byte, 1<<20) + n := runtime.Stack(stack, true) + if n == len(stack) { + t.Error("owned Close stack observation was truncated") + return false + } + for _, block := range strings.Split(string(stack[:n]), "\n\n") { + line, _, _ := strings.Cut(block, "\n") + if strings.Contains(line, "[chan receive]") && + strings.Contains(block, "(*WebH2Client).Close(") && + strings.Contains(block, "webH2CloseJoinCall(") { + // In these fixtures the first Close has no registered socket to + // close: the pending dial or detached cleanup is outside the map. + // Observe the already-closed caller's actual completion wait. + t.Log("observed own concurrent Close at the shared completion wait") + select { + case <-first: + t.Error("first Close escaped pending cleanup") + return false + case <-second: + t.Error("concurrent Close escaped pending cleanup") + return false + case <-time.After(75 * time.Millisecond): + return true + } + } + } + select { + case <-first: + t.Error("first Close returned before cleanup join") + return false + case <-second: + t.Error("concurrent Close returned before cleanup join") + return false + case <-timer.C: + t.Error("concurrent Close did not reach an actual completion wait") + return false + case <-ticker.C: + } + } +} + +func webH2CloseJoinStartClosers(t *testing.T, c *WebH2Client, results chan<- error, workers *[]<-chan struct{}) (<-chan struct{}, <-chan struct{}) { + t.Helper() + started := make(chan struct{}, 2) + first, second := make(chan struct{}), make(chan struct{}) + *workers = append(*workers, first, second) + go webH2CloseJoinCall(c, started, results, first) + webH3CloseJoinRead(t, started, "first H2 Close entry") + select { + case <-c.ctx.Done(): + case <-time.After(time.Second): + t.Fatal("first H2 Close did not cancel its actual client lifetime") + } + go webH2CloseJoinCall(c, started, results, second) + webH3CloseJoinRead(t, started, "second H2 Close entry") + return first, second +} + +func webH2CloseJoinAssertFinal(t *testing.T, c *WebH2Client) { + t.Helper() + c.mu.Lock() + closed, count, current := c.closed, len(c.sessions), c.current + c.mu.Unlock() + if !closed || count != 0 || current != nil || len(c.dialGate) != 0 { + t.Errorf("joined Close state closed/sessions/current/gate=%t/%d/%p/%d", closed, count, current, len(c.dialGate)) + } + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := c.DialContext(ctx, "tcp", "target.invalid:443") + if conn != nil { + _ = conn.Close() + t.Error("closed client returned a new tunnel connection") + } + if !errors.Is(err, net.ErrClosed) { + t.Errorf("post-Close dial=%v, want net.ErrClosed", err) + } +} + +func TestWebH2CloseJoinsLatePhysicalDial(t *testing.T) { + for _, tc := range []struct { + name string + conn bool + err error + }{ + {name: "owned_conn_success", conn: true}, + {name: "owned_conn_error", conn: true, err: errors.New("private late physical dial error")}, + {name: "nil_conn_error", err: errors.New("private late nil dial error")}, + } { + t.Run(tc.name, func(t *testing.T) { + listener, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1)}) + if err != nil { + t.Fatal(err) + } + if err := listener.SetDeadline(time.Now().Add(2 * time.Second)); err != nil { + _ = listener.Close() + t.Fatal(err) + } + _, clientTLS := testTLSConfigs(t) + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: listener.Addr().String(), Token: webTestToken, TLSConfig: clientTLS, + FingerprintProfile: FingerprintNative, HandshakeTimeout: time.Second, DialTimeout: time.Second, + }) + if err != nil { + _ = listener.Close() + t.Fatal(err) + } + fixtureCtx, stopFixture := context.WithTimeout(context.Background(), 4*time.Second) + defer stopFixture() + releaseDial, releaseClose := make(chan struct{}), make(chan struct{}) + ungateDial := sync.OnceFunc(func() { close(releaseDial) }) + ungateClose := sync.OnceFunc(func() { close(releaseClose) }) + physicalCtx := make(chan context.Context, 1) + physicalReturned := make(chan struct{}) + physicalReturnOnce := sync.OnceFunc(func() { close(physicalReturned) }) + var owned atomic.Pointer[webH2CloseJoinHeldConn] + var dials atomic.Int32 + client.dialer = transport.DialFunc(func(ctx context.Context, network, address string) (net.Conn, error) { + defer physicalReturnOnce() + if dials.Add(1) != 1 { + return nil, errors.New("late-dial fixture observed an unexpected extra attempt") + } + if network != "tcp" || address != listener.Addr().String() { + return nil, errors.New("late-dial fixture refuses a non-owned endpoint") + } + var wire *webH2CloseJoinHeldConn + if tc.conn { + raw, dialErr := (&net.Dialer{Timeout: time.Second}).DialContext(ctx, network, address) + if dialErr != nil { + return nil, dialErr + } + wire = &webH2CloseJoinHeldConn{Conn: raw, entered: make(chan struct{}), returned: make(chan struct{}), release: releaseClose} + owned.Store(wire) + } + physicalCtx <- ctx + // A real dial can complete concurrently with lifetime cancellation. + // The separate bounded fixture gate controls its late publication. + select { + case <-releaseDial: + case <-fixtureCtx.Done(): + } + if wire == nil { + return nil, tc.err + } + return wire, tc.err + }) + var peer *net.TCPConn + var closeWorkers []<-chan struct{} + dialDone, dialResult := make(chan struct{}), make(chan error, 1) + t.Cleanup(func() { + ungateDial() + ungateClose() + stopFixture() + if wire := owned.Load(); wire != nil { + _ = wire.Conn.Close() // Independent exact-socket cleanup, not an oracle. + } + if peer != nil { + _ = peer.Close() + } + _ = listener.Close() + webH3CloseJoinWait(t, dialDone, "late public H2 dial worker cleanup") + for _, worker := range closeWorkers { + webH3CloseJoinWait(t, worker, "H2 Close caller cleanup") + } + cleanupDone := make(chan struct{}) + go func() { defer close(cleanupDone); _ = client.Close() }() + webH3CloseJoinWait(t, cleanupDone, "independent H2 client cleanup") + }) + go func() { + defer close(dialDone) + conn, err := client.DialContext(fixtureCtx, "tcp", "target.invalid:443") + if conn != nil { + _ = conn.Close() + err = errors.New("closed client returned a late tunnel connection") + } + dialResult <- err + }() + ctx := webH3CloseJoinRead(t, physicalCtx, "actual owned physical dial entry") + if tc.conn { + peer, err = listener.AcceptTCP() + if err != nil { + t.Fatal(err) + } + } + results := make(chan error, 2) + first, second := webH2CloseJoinStartClosers(t, client, results, &closeWorkers) + select { + case <-ctx.Done(): + case <-time.After(time.Second): + t.Error("Close did not cancel the in-flight physical dial") + } + webH2CloseJoinCheckPending(t, first, second) + ungateDial() + webH3CloseJoinMustWait(t, physicalReturned, "actual late physical dial return") + if wire := owned.Load(); wire != nil { + select { + case <-wire.entered: + webH2CloseJoinCheckPending(t, first, second) + case <-time.After(time.Second): + t.Error("late owned physical connection was not closed by the client") + } + if err := peer.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil { + t.Fatal(err) + } + var one [1]byte + n, readErr := peer.Read(one[:]) + if n != 0 || !(errors.Is(readErr, io.EOF) || errors.Is(readErr, syscall.ECONNRESET)) { + t.Errorf("late owned socket not actually closed before independent cleanup: n=%d error=%v", n, readErr) + } + } + ungateClose() + webH3CloseJoinMustWait(t, first, "first H2 Close result") + webH3CloseJoinMustWait(t, second, "concurrent H2 Close result") + webH3CloseJoinMustWait(t, dialDone, "late public H2 dial completion") + if err := webH3CloseJoinRead(t, dialResult, "late public H2 dial result"); !errors.Is(err, net.ErrClosed) { + t.Errorf("dial after client Close=%v, want net.ErrClosed", err) + } + for range 2 { + if err := webH3CloseJoinRead(t, results, "joined H2 Close result"); err != nil { + t.Errorf("joined H2 Close=%v", err) + } + } + webH2CloseJoinAssertFinal(t, client) + if dials.Load() != 1 { + t.Errorf("physical dial count=%d, want one exact owned attempt", dials.Load()) + } + }) + } +} + +func TestWebH2CloseJoinsDetachedSessionCleanup(t *testing.T) { + payload := []byte("owned H2 detached cleanup with complete upload FIN") + origin := newWebH3CloseJoinOrigin(t, payload) // Bounded actual TCP helper, not H3 setup. + target := origin.listener.Addr().String() + serverTLS, clientTLS := testTLSConfigs(t) + server, err := ListenWebH2(WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: webTestToken, TLSConfig: serverTLS, Cover: http.NotFoundHandler(), + Dialer: transport.DialFunc(func(ctx context.Context, network, address string) (net.Conn, error) { + if network != "tcp" || address != target { + return nil, errors.New("warm cleanup fixture refuses a non-owned target") + } + return (&net.Dialer{Timeout: time.Second}).DialContext(ctx, network, address) + }), + }) + if err != nil { + t.Fatal(err) + } + handler := server.server.Handler + handlerDone, bearerLength := make(chan struct{}), make(chan int, 1) + handlerReturned := sync.OnceFunc(func() { close(handlerDone) }) + var requests, unexpected atomic.Int32 + server.server.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer handlerReturned() + requests.Add(1) + select { + case bearerLength <- len(r.Header.Get("Proxy-Authorization")): + default: + unexpected.Add(1) + } + handler.ServeHTTP(w, r) + }) + serveCtx, cancelServe := context.WithCancel(context.Background()) + serveDone, serveResult := make(chan struct{}), make(chan error, 1) + go func() { defer close(serveDone); serveResult <- server.Serve(serveCtx) }() + t.Cleanup(func() { + _ = server.Close() + cancelServe() + webH3CloseJoinWait(t, serveDone, "owned H2 Serve cleanup") + }) + client, err := NewWebH2Client(WebH2ClientConfig{ServerAddress: server.Addr().String(), Token: webTestToken, TLSConfig: clientTLS, FingerprintProfile: FingerprintNative}) + if err != nil { + t.Fatal(err) + } + release := make(chan struct{}) + ungate := sync.OnceFunc(func() { close(release) }) + var held *webH2CloseJoinHeldConn + var stream net.Conn + var closeWorkers []<-chan struct{} + retiredDone := make(chan struct{}) + retiredStarted := false + t.Cleanup(func() { + ungate() + if held != nil { + _ = held.Conn.Close() + } + if stream != nil { + _ = stream.Close() + } + cleanupDone := make(chan struct{}) + go func() { defer close(cleanupDone); _ = client.Close() }() + webH3CloseJoinWait(t, cleanupDone, "independent warm H2 client cleanup") + for _, worker := range closeWorkers { + webH3CloseJoinWait(t, worker, "warm H2 Close caller cleanup") + } + if retiredStarted { + webH3CloseJoinWait(t, retiredDone, "detached H2 retirement cleanup") + } + }) + dialCtx, cancelDial := context.WithTimeout(context.Background(), 2*time.Second) + defer cancelDial() + stream, err = client.DialContext(dialCtx, "tcp", target) + if err != nil { + t.Fatal(err) + } + if err := stream.SetDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + if _, err := stream.Write(payload); err != nil { + t.Fatal(err) + } + if err := stream.(interface{ CloseWrite() error }).CloseWrite(); err != nil { + t.Fatal(err) + } + reply, err := io.ReadAll(stream) + if err != nil || !bytes.Equal(reply, append([]byte("reply:"), payload...)) { + t.Fatalf("actual healthy H2 payload/reply=%q/%v", reply, err) + } + webH3CloseJoinMustWait(t, origin.done, "owned complete TCP origin") + if err := webH3CloseJoinRead(t, origin.result, "actual origin result"); err != nil { + t.Fatal(err) + } + webH3CloseJoinMustWait(t, handlerDone, "actual authenticated H2 CONNECT handler") + if length := webH3CloseJoinRead(t, bearerLength, "actual full H2 bootstrap"); length < 385 || length > 2047 { + t.Fatalf("full bootstrap credential length=%d", length) + } + if requests.Load() != 1 || unexpected.Load() != 0 { + t.Fatalf("actual warm request/observer counts=%d/%d", requests.Load(), unexpected.Load()) + } + if err := stream.Close(); err != nil { + t.Fatal(err) + } + client.mu.Lock() + session := client.current + if session == nil || len(client.sessions) != 1 || session.authState != webH2ClientAuthReady || session.auth == nil || session.opening != 0 || session.active != 0 || !session.h2.CanTakeNewRequest() { + client.mu.Unlock() + t.Fatal("exact live authenticated idle H2 session not proved before instrumentation") + } + held = &webH2CloseJoinHeldConn{Conn: session.raw, entered: make(chan struct{}), returned: make(chan struct{}), release: release} + session.raw = held + client.mu.Unlock() + retiredStarted = true + go func() { + defer close(retiredDone) + session.h2.SetDoNotReuse() + client.noteSessionFailure(session) + }() + webH3CloseJoinMustWait(t, held.entered, "actual raw socket closed before detached completion gate") + client.mu.Lock() + removed := client.current == nil && len(client.sessions) == 0 + client.mu.Unlock() + if !removed { + t.Fatal("retirement did not remove the exact session before detached cleanup") + } + results := make(chan error, 2) + first, second := webH2CloseJoinStartClosers(t, client, results, &closeWorkers) + webH2CloseJoinCheckPending(t, first, second) + ungate() + webH3CloseJoinMustWait(t, retiredDone, "actual detached H2 cleanup completion") + webH3CloseJoinMustWait(t, held.returned, "actual held TCP Close completion") + webH3CloseJoinMustWait(t, first, "first warm H2 Close") + webH3CloseJoinMustWait(t, second, "concurrent warm H2 Close") + for range 2 { + if err := webH3CloseJoinRead(t, results, "joined warm H2 Close result"); err != nil { + t.Errorf("joined warm H2 Close=%v", err) + } + } + webH2CloseJoinAssertFinal(t, client) + _ = server.Close() + cancelServe() + webH3CloseJoinMustWait(t, serveDone, "normal H2 Serve shutdown") + if err := webH3CloseJoinRead(t, serveResult, "actual H2 Serve result"); err != nil { + t.Errorf("H2 Serve=%v", err) + } +} diff --git a/internal/tunnel/web_h2_dial_result_test.go b/internal/tunnel/web_h2_dial_result_test.go new file mode 100644 index 0000000..eb7efd7 --- /dev/null +++ b/internal/tunnel/web_h2_dial_result_test.go @@ -0,0 +1,114 @@ +package tunnel + +import ( + "context" + "errors" + "net" + "sync/atomic" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +// Models a context deadline whose timer has not yet published Done. The +// caller's fixed deadline is authoritative even if its timer loses the race. +type webH2UnpublishedDeadlineContext struct { + context.Context + deadline time.Time +} + +func (c webH2UnpublishedDeadlineContext) Deadline() (time.Time, bool) { + return c.deadline, true +} + +func TestWebH2DialContextPreflightCause(t *testing.T) { + for _, mode := range []string{"canceled_cause", "elapsed_unpublished_deadline"} { + t.Run(mode, func(t *testing.T) { + _, clientTLS := testTLSConfigs(t) + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: "127.0.0.1:443", Token: testToken, TLSConfig: clientTLS, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + var calls atomic.Int32 + client.dialer = transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + calls.Add(1) + return nil, errors.New("expired caller reached physical dial") + }) + var ctx context.Context + want := error(context.DeadlineExceeded) + if mode == "canceled_cause" { + cause := errors.New("private H2 caller cancellation") + parent, cancel := context.WithCancelCause(context.Background()) + cancel(cause) + ctx, want = parent, cause + } else { + ctx = webH2UnpublishedDeadlineContext{Context: context.Background(), deadline: time.Now().Add(-time.Second)} + } + conn, err := client.DialContext(ctx, "tcp", "target.invalid:443") + if conn != nil { + _ = conn.Close() + t.Error("expired caller received a connection") + } + if !errors.Is(err, want) || calls.Load() != 0 { + t.Errorf("error=%v physical dials=%d; want cause %v and no dial", err, calls.Load(), want) + } + }) + } +} + +func TestWebH2DialConnectionErrorOwnership(t *testing.T) { + _, clientTLS := testTLSConfigs(t) + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: "127.0.0.1:443", Token: testToken, TLSConfig: clientTLS, + }) + if err != nil { + t.Fatal(err) + } + raw, peer := net.Pipe() + wire := &webH2InitializationWire{Conn: raw, closed: make(chan struct{})} + t.Cleanup(func() { _ = raw.Close(); _ = peer.Close(); _ = client.Close() }) + fault := errors.New("physical TCP dial returned connection and failure") + client.dialer = transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { return wire, fault }) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := client.DialContext(ctx, "tcp", "target.invalid:443") + if conn != nil { + _ = conn.Close() + t.Error("failed physical dial returned a stream") + } + if !errors.Is(err, fault) { + t.Errorf("dial error=%v, want original failure", err) + } + select { + case <-wire.closed: + default: + t.Error("failed physical dial retained its non-nil connection") + } +} + +// Old production panics for nil+nil; this guard is deliberately excluded from +// old-source negative runs, rather than treating a process crash as evidence. +func TestWebH2DialMissingConnectionFailsClosed(t *testing.T) { + _, clientTLS := testTLSConfigs(t) + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: "127.0.0.1:443", Token: testToken, TLSConfig: clientTLS, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + client.dialer = transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { return nil, nil }) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + conn, err := client.DialContext(ctx, "tcp", "target.invalid:443") + if conn != nil || err == nil { + if conn != nil { + _ = conn.Close() + } + t.Fatalf("nil physical result: conn=%v error=%v", conn, err) + } +} diff --git a/internal/tunnel/web_h2_dial_sharing_test.go b/internal/tunnel/web_h2_dial_sharing_test.go new file mode 100644 index 0000000..483380a --- /dev/null +++ b/internal/tunnel/web_h2_dial_sharing_test.go @@ -0,0 +1,706 @@ +package tunnel + +import ( + "context" + "crypto/tls" + "encoding/binary" + "errors" + "fmt" + "io" + "net" + "net/http" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +// All networking is owned IPv4 loopback. Independent fixture cancellation, +// actual socket closure, and separate goroutine-completion channels bound +// cleanup, including assertion failures against the old implementation. +type webH2DialSharingFixture struct { + t *testing.T + ctx context.Context + cancel context.CancelFunc + mu sync.Mutex + joins []<-chan struct{} + close []func() + checks []func() +} + +func newWebH2DialSharingFixture(t *testing.T) *webH2DialSharingFixture { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 6*time.Second) + f := &webH2DialSharingFixture{t: t, ctx: ctx, cancel: cancel} + t.Cleanup(func() { + // Close the server before canceling Serve's context to avoid racing two + // independently triggered server Close paths in this fixture. + for i := len(f.close) - 1; i >= 0; i-- { + f.close[i]() + } + cancel() + timer := time.NewTimer(2 * time.Second) + defer timer.Stop() + // Every parent is registered before launch; it registers its children + // before returning. Join parents before taking each fresh next entry, + // so late accept/forward/handler registration cannot escape cleanup. + for i := 0; ; i++ { + f.mu.Lock() + if i == len(f.joins) { + f.mu.Unlock() + break + } + joined := f.joins[i] + f.mu.Unlock() + select { + case <-joined: + case <-timer.C: + t.Errorf("cleanup did not join owned worker %d", i) + return + } + } + for _, check := range f.checks { + check() + } + }) + return f +} + +func (f *webH2DialSharingFixture) register(done <-chan struct{}) { + f.mu.Lock() + f.joins = append(f.joins, done) + f.mu.Unlock() +} + +func (f *webH2DialSharingFixture) wait(done <-chan struct{}, name string) { + f.t.Helper() + select { + case <-done: + case <-f.ctx.Done(): + f.t.Fatalf("%s exceeded independent fixture budget", name) + } +} + +type webH2DialSharingContext struct { + context.Context + entered chan struct{} + attemptWait chan struct{} + doneEvaluated atomic.Int32 +} + +// First Done evaluation only witnesses the selection gate. The second is the +// wait select after a brief gate's attempt capture. Both must be observed +// before the real peer is released. Old code blocks in the first gate select, +// so the bounded checkpoint fails honestly before testing its extra dials. +// Public AfterFunc registration is deliberately not used as a join witness. +func (c *webH2DialSharingContext) Done() <-chan struct{} { + switch c.doneEvaluated.Add(1) { + case 1: + close(c.entered) + case 2: + close(c.attemptWait) + } + return c.Context.Done() +} + +func (f *webH2DialSharingFixture) joinedWaiters(witnesses []*webH2DialSharingContext) { + f.t.Helper() + deadline := time.NewTimer(500 * time.Millisecond) + defer deadline.Stop() + for i, witness := range witnesses { + select { + case <-witness.attemptWait: + case <-deadline.C: + f.t.Errorf("follower %d did not reach shared-attempt wait after selection gate; Done evaluations=%d", i, witness.doneEvaluated.Load()) + return + case <-f.ctx.Done(): + f.t.Fatal("independent fixture expired at joined-attempt checkpoint") + } + } +} + +type webH2DialSharingRaw struct { + net.Conn + closed chan struct{} + once sync.Once +} + +func (c *webH2DialSharingRaw) Close() error { + err := c.Conn.Close() + c.once.Do(func() { close(c.closed) }) + return err +} + +type webH2DialSharingFront struct { + f *webH2DialSharingFixture + listener net.Listener + destination string + release chan struct{} + firstHello chan struct{} + releaseOnce sync.Once + helloOnce sync.Once + mu sync.Mutex + conns []net.Conn + accepts int + hellos int + err error + closed bool +} + +// With no destination this is a real TCP/TLS blackhole until release, after +// which the peer closes each actual ClientHello connection. Otherwise it +// forwards the captured bytes unchanged to an owned real TLS/H2 server. +// No TLS/HTTP responses, canceled context, or reader stalls are synthesized. +func (f *webH2DialSharingFixture) front(destination string) *webH2DialSharingFront { + f.t.Helper() + listener, err := net.Listen("tcp4", "127.0.0.1:0") + if err != nil { + f.t.Fatal(err) + } + if destination != "" { + host, _, err := net.SplitHostPort(destination) + if err != nil || host != "127.0.0.1" { + f.t.Fatalf("non-owned forwarding destination %q", destination) + } + } + p := &webH2DialSharingFront{f: f, listener: listener, destination: destination, + release: make(chan struct{}), firstHello: make(chan struct{})} + f.close = append(f.close, p.Close) + joined := make(chan struct{}) + f.register(joined) + go func() { + defer close(joined) + for { + conn, err := listener.Accept() + if err != nil { + p.noteError(err) + return + } + p.mu.Lock() + if p.closed { + p.mu.Unlock() + _ = conn.Close() + continue + } + p.accepts++ + p.conns = append(p.conns, conn) + p.mu.Unlock() + _ = conn.SetDeadline(time.Now().Add(7 * time.Second)) + workerDone := make(chan struct{}) + f.register(workerDone) + go func() { + defer close(workerDone) + defer conn.Close() + p.serve(conn) + }() + } + }() + return p +} + +func (p *webH2DialSharingFront) noteError(err error) { + if err == nil || errors.Is(err, net.ErrClosed) || errors.Is(err, io.EOF) || errors.Is(err, context.Canceled) { + return + } + p.mu.Lock() + if !p.closed && p.f.ctx.Err() == nil { + p.err = errors.Join(p.err, err) + } + p.mu.Unlock() +} + +func (p *webH2DialSharingFront) serve(front net.Conn) { + var header [5]byte + if _, err := io.ReadFull(front, header[:]); err != nil { + p.noteError(err) + return + } + n := int(binary.BigEndian.Uint16(header[3:])) + if header[0] != 22 || n < 4 || n > 16<<10 { + p.noteError(fmt.Errorf("first wire record is not a bounded TLS handshake: type=%d bytes=%d", header[0], n)) + return + } + hello := make([]byte, 5+n) + copy(hello, header[:]) + if _, err := io.ReadFull(front, hello[5:]); err != nil { + p.noteError(err) + return + } + if hello[5] != 1 { // TLS handshake message ClientHello. + p.noteError(errors.New("first actual TLS message was not ClientHello")) + return + } + p.mu.Lock() + p.hellos++ + p.mu.Unlock() + p.helloOnce.Do(func() { close(p.firstHello) }) + select { + case <-p.release: + case <-p.f.ctx.Done(): + return + } + if p.destination == "" { + return // Deferred actual peer Close, not a synthetic dial error. + } + back, err := (&net.Dialer{}).DialContext(p.f.ctx, "tcp4", p.destination) + if err != nil { + p.noteError(err) + return + } + defer back.Close() + _ = back.SetDeadline(time.Now().Add(7 * time.Second)) + p.mu.Lock() + if p.closed { + p.mu.Unlock() + return + } + p.conns = append(p.conns, back) + p.mu.Unlock() + if _, err := back.Write(hello); err != nil { + p.noteError(err) + return + } + result := make(chan error, 2) + for _, pair := range [][2]net.Conn{{back, front}, {front, back}} { + done := make(chan struct{}) + p.f.register(done) + go func(destination, source net.Conn) { + defer close(done) + _, err := io.Copy(destination, source) + if tcp, ok := destination.(*net.TCPConn); ok { + _ = tcp.CloseWrite() + } + result <- err + }(pair[0], pair[1]) + } + for range 2 { + if err := <-result; err != nil { + // Peer cancellation/reset is expected when running the old policy. + _ = front.Close() + _ = back.Close() + } + } +} + +func (p *webH2DialSharingFront) Release() { p.releaseOnce.Do(func() { close(p.release) }) } + +func (p *webH2DialSharingFront) Close() { + p.mu.Lock() + p.closed = true + conns := append([]net.Conn(nil), p.conns...) + p.mu.Unlock() + p.Release() + _ = p.listener.Close() + for _, conn := range conns { + _ = conn.Close() + } +} + +func (p *webH2DialSharingFront) state() (int, int, error) { + p.mu.Lock() + defer p.mu.Unlock() + return p.accepts, p.hellos, p.err +} + +func (f *webH2DialSharingFixture) client(profile FingerprintProfile, address string, config *tls.Config) (*WebH2Client, *webH2AuthEntropyCounter, <-chan *webH2DialSharingRaw) { + f.t.Helper() + entropy := &webH2AuthEntropyCounter{} + client, err := newWebH2ClientWithSigner(WebH2ClientConfig{ + ServerAddress: address, Token: webTestToken, TLSConfig: config, FingerprintProfile: profile, + DialTimeout: time.Second, HandshakeTimeout: 2 * time.Second, + }, newWebAuthSigner(mustWebAuthKey(f.t, webTestToken), nil, entropy), webAuthClaims{}) + if err != nil { + f.t.Fatal(err) + } + raws := make(chan *webH2DialSharingRaw, 8) + client.dialer = transport.DialFunc(func(ctx context.Context, network, target string) (net.Conn, error) { + if network != "tcp" || target != address { + return nil, errors.New("non-owned physical dial") + } + raw, err := (&net.Dialer{Timeout: time.Second}).DialContext(ctx, "tcp4", address) + if err != nil { + return nil, err + } + wire := &webH2DialSharingRaw{Conn: raw, closed: make(chan struct{})} + raws <- wire + return wire, nil + }) + f.close = append(f.close, func() { + if err := client.Close(); err != nil { + f.t.Errorf("owned client Close: %v", err) + } + }) + return client, entropy, raws +} + +type webH2DialSharingResult struct { + conn net.Conn + err error +} + +func (f *webH2DialSharingFixture) reserve(client *WebH2Client, ctx context.Context) (<-chan webH2DialSharingResult, <-chan struct{}) { + results, joined := make(chan webH2DialSharingResult, 1), make(chan struct{}) + f.register(joined) + go func() { + defer close(joined) + reservation, err := client.reserveSession(ctx) + if reservation.session != nil { + client.releaseSessionReservation(reservation.session) + } + results <- webH2DialSharingResult{err: err} + }() + return results, joined +} + +func (f *webH2DialSharingFixture) dial(client *WebH2Client, ctx context.Context, target string) (<-chan webH2DialSharingResult, <-chan struct{}) { + results, joined := make(chan webH2DialSharingResult, 1), make(chan struct{}) + f.register(joined) + go func() { + defer close(joined) + conn, err := client.DialContext(ctx, "tcp", target) + results <- webH2DialSharingResult{conn: conn, err: err} + }() + return results, joined +} + +func (f *webH2DialSharingFixture) raw(raws <-chan *webH2DialSharingRaw) *webH2DialSharingRaw { + f.t.Helper() + select { + case raw := <-raws: + return raw + case <-f.ctx.Done(): + f.t.Fatal("actual physical dial did not return an owned raw connection") + return nil + } +} + +func TestWebH2DialSharingFailureAndLaterRetry(t *testing.T) { + for _, profile := range []FingerprintProfile{FingerprintNative, FingerprintChrome133} { + t.Run(string(profile), func(t *testing.T) { + f := newWebH2DialSharingFixture(t) + _, clientTLS := testTLSConfigs(t) + front := f.front("") + client, entropy, raws := f.client(profile, front.listener.Addr().String(), clientTLS) + results := make([]<-chan webH2DialSharingResult, 3) + joins := make([]<-chan struct{}, 3) + results[0], joins[0] = f.reserve(client, f.ctx) + f.wait(front.firstHello, "real first ClientHello") + firstRaw := f.raw(raws) + var witnesses []*webH2DialSharingContext + for i := 1; i < 3; i++ { + witness := &webH2DialSharingContext{Context: f.ctx, entered: make(chan struct{}), attemptWait: make(chan struct{})} + results[i], joins[i] = f.reserve(client, witness) + f.wait(witness.entered, "direct reserve follower wait selection") + witnesses = append(witnesses, witness) + } + f.joinedWaiters(witnesses) + if accepts, hellos, err := front.state(); accepts != 1 || hellos != 1 || err != nil { + t.Fatalf("pre-failure physical accepts/hellos/error=%d/%d/%v", accepts, hellos, err) + } + front.Release() + var firstError error + for i := range results { + f.wait(joins[i], "failed reserve caller joined") + err := (<-results[i]).err + if err == nil { + t.Error("real closed TLS peer unexpectedly initialized a session") + } + if i == 0 { + firstError = err + } else if err != firstError { + t.Errorf("joined caller %d received a new failure object, not the shared immutable result", i) + } + } + f.wait(firstRaw.closed, "failed physical raw Close completion") + if accepts, hellos, err := front.state(); accepts != 1 || hellos != 1 || err != nil { + t.Errorf("one joined failed attempt produced physical accepts/hellos/error=%d/%d/%v, want 1/1/nil", accepts, hellos, err) + } + later, laterJoined := f.reserve(client, f.ctx) + f.wait(laterJoined, "later independent retry joined") + if err := (<-later).err; err == nil || err == firstError { + t.Errorf("later retry must be a new failed attempt/result: %v", err) + } + if accepts, hellos, err := front.state(); accepts != 2 || hellos != 2 || err != nil { + t.Errorf("later retry physical accepts/hellos/error=%d/%d/%v, want 2/2/nil", accepts, hellos, err) + } + if entropy.nonceReads.Load() != 0 { + t.Error("failed physical initialization generated application authentication") + } + t.Logf("actual TLS ClientHello attempts and same-result/later-retry oracles evaluated; nonce reads=%d", entropy.nonceReads.Load()) + }) + } +} + +func (f *webH2DialSharingFixture) target() string { + f.t.Helper() + listener, err := net.Listen("tcp4", "127.0.0.1:0") + if err != nil { + f.t.Fatal(err) + } + var mu sync.Mutex + var conns []net.Conn + closed := false + f.close = append(f.close, func() { + mu.Lock() + closed = true + owned := append([]net.Conn(nil), conns...) + mu.Unlock() + _ = listener.Close() + for _, conn := range owned { + _ = conn.Close() + } + }) + joined := make(chan struct{}) + f.register(joined) + go func() { + defer close(joined) + for { + conn, err := listener.Accept() + if err != nil { + return + } + mu.Lock() + if closed { + mu.Unlock() + _ = conn.Close() + return + } + conns = append(conns, conn) + mu.Unlock() + _ = conn.SetDeadline(time.Now().Add(7 * time.Second)) + done := make(chan struct{}) + f.register(done) + go func() { + defer close(done) + defer conn.Close() + var payload [256]byte + for { + n, err := conn.Read(payload[:]) + if n != 0 { + if _, writeErr := conn.Write(payload[:n]); writeErr != nil { + return + } + } + if err != nil { + return + } + } + }() + } + }() + return listener.Addr().String() +} + +func TestWebH2DialInitiatingCallerCancellationKeepsPhysicalInitialization(t *testing.T) { + for _, profile := range []FingerprintProfile{FingerprintNative, FingerprintChrome133} { + t.Run(string(profile), func(t *testing.T) { + f := newWebH2DialSharingFixture(t) + serverTLS, clientTLS := testTLSConfigs(t) + target := f.target() + var targetDials, covers atomic.Int64 + server, err := ListenWebH2(WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: webTestToken, TLSConfig: serverTLS, + Cover: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { covers.Add(1); http.NotFound(w, r) }), + Dialer: transport.DialFunc(func(ctx context.Context, network, address string) (net.Conn, error) { + targetDials.Add(1) + if network != "tcp" || address != target { + return nil, errors.New("non-owned destination dial") + } + return (&net.Dialer{}).DialContext(ctx, "tcp4", target) + }), + }) + if err != nil { + t.Fatal(err) + } + var requestsMu sync.Mutex + var authConnections []*webServerConnectionAuth + var bearerLengths []int + original := server.server.Handler + server.server.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + handlerJoined := make(chan struct{}) + f.register(handlerJoined) + defer close(handlerJoined) + state, _ := r.Context().Value(webServerConnectionAuthContextKey{}).(*webServerConnectionAuth) + requestsMu.Lock() + authConnections = append(authConnections, state) + bearerLengths = append(bearerLengths, len(r.Header.Get("Proxy-Authorization"))) + requestsMu.Unlock() + original.ServeHTTP(w, r) + }) + serveJoined := make(chan struct{}) + var serveErr error + f.register(serveJoined) + f.checks = append(f.checks, func() { + if serveErr != nil { + t.Errorf("owned H2 Serve result: %v", serveErr) + } + }) + f.close = append(f.close, func() { + if err := server.Close(); err != nil { + t.Errorf("owned H2 server Close: %v", err) + } + }) + go func() { defer close(serveJoined); serveErr = server.Serve(f.ctx) }() + front := f.front(server.Addr().String()) + client, entropy, raws := f.client(profile, front.listener.Addr().String(), clientTLS) + leaderCtx, cancelLeader := context.WithCancelCause(f.ctx) + defer cancelLeader(context.Canceled) + leader, leaderJoined := f.dial(client, leaderCtx, target) + f.wait(front.firstHello, "actual gated public leader ClientHello") + firstRaw := f.raw(raws) + cause := errors.New("initiating H2 caller abandoned its wait") + cancelLeader(cause) + f.wait(leaderJoined, "canceled public initiating caller joined") + result := <-leader + if result.conn != nil { + _ = result.conn.Close() + t.Error("canceled initiating caller returned an established stream") + } + if !errors.Is(result.err, cause) { + t.Errorf("initiating caller error=%v, want its custom cause", result.err) + } + select { + case <-firstRaw.closed: + t.Error("initiating caller cancellation closed the shared cold physical raw connection") + default: + } + front.Release() + // No live caller exists now. A client-owned successful initializer + // must publish an unclaimed physical session, not send CONNECT/auth. + var initialized *webH2ClientSession + ticker, deadline := time.NewTicker(10*time.Millisecond), time.NewTimer(time.Second) + defer ticker.Stop() + defer deadline.Stop() + waitInitialized: + for { + client.mu.Lock() + initialized = client.current + unclaimed := initialized != nil && initialized.auth == nil && initialized.opening == 0 + client.mu.Unlock() + if unclaimed { + break + } + select { + case <-ticker.C: + case <-deadline.C: + t.Error("all waits abandoned: physical initialization did not publish an unclaimed session") + break waitInitialized + case <-f.ctx.Done(): + t.Fatal("independent fixture expired before physical initialization") + } + } + requestsMu.Lock() + before := len(authConnections) + requestsMu.Unlock() + if before != 0 || targetDials.Load() != 0 || entropy.nonceReads.Load() != 0 { + t.Errorf("unclaimed physical session sent requests/target/auth=%d/%d/%d", before, targetDials.Load(), entropy.nonceReads.Load()) + } + for round := range 2 { + caller, cancel := context.WithCancel(f.ctx) + conn, err := client.DialContext(caller, "tcp", target) + cancel() // A returned stream is not owned by its establishment context. + if err != nil { + t.Fatalf("live caller %d real CONNECT: %v", round, err) + } + _ = conn.SetDeadline(time.Now().Add(time.Second)) + payload := []byte(fmt.Sprintf("owned-h2-survivor-%d", round)) + _, writeErr := conn.Write(payload) + got := make([]byte, len(payload)) + _, readErr := io.ReadFull(conn, got) + _ = conn.Close() + if writeErr != nil || readErr != nil || string(got) != string(payload) { + t.Fatalf("live caller %d full echo=%q, write/read=%v/%v", round, got, writeErr, readErr) + } + } + client.mu.Lock() + warm := client.current + client.mu.Unlock() + if initialized != nil && warm != initialized { + t.Error("live caller did not reuse the initialized unclaimed physical session") + } + requestsMu.Lock() + samePhysical := len(authConnections) == 2 && authConnections[0] != nil && authConnections[0] == authConnections[1] + fullShort := len(bearerLengths) == 2 && bearerLengths[0] > bearerLengths[1] && bearerLengths[1] != 0 + requestsMu.Unlock() + if !samePhysical || !fullShort || entropy.nonceReads.Load() != 1 || targetDials.Load() != 2 || covers.Load() != 0 { + t.Errorf("healthy full/short auth physical=%t fullShort=%t nonce/target/cover=%d/%d/%d", samePhysical, fullShort, entropy.nonceReads.Load(), targetDials.Load(), covers.Load()) + } + if accepts, hellos, err := front.state(); accepts != 1 || hellos != 1 || err != nil { + t.Errorf("caller cancellation caused replacement physical accepts/hellos/error=%d/%d/%v, want 1/1/nil", accepts, hellos, err) + } + if err := client.Close(); err != nil { + t.Error(err) + } + if err := server.Close(); err != nil { + t.Error(err) + } + f.wait(serveJoined, "owned H2 Serve joined") + t.Log("actual public survivor and warm sibling full echoes succeeded; one full auth plus short continuation evaluated") + }) + } +} + +func TestWebH2DialAllWaitersAbandonThenClose(t *testing.T) { + for _, profile := range []FingerprintProfile{FingerprintNative, FingerprintChrome133} { + t.Run(string(profile), func(t *testing.T) { + f := newWebH2DialSharingFixture(t) + _, clientTLS := testTLSConfigs(t) + front := f.front("") + client, entropy, raws := f.client(profile, front.listener.Addr().String(), clientTLS) + results := make([]<-chan webH2DialSharingResult, 3) + joins := make([]<-chan struct{}, 3) + cancels := make([]context.CancelCauseFunc, 3) + causes := make([]error, 3) + var witnesses []*webH2DialSharingContext + for i := range 3 { + ctx, cancel := context.WithCancelCause(f.ctx) + cancels[i] = cancel + defer cancel(context.Canceled) + causes[i] = fmt.Errorf("H2 waiter %d abandoned", i) + if i == 0 { + results[i], joins[i] = f.reserve(client, ctx) + f.wait(front.firstHello, "actual first ClientHello before abandonment") + } else { + witness := &webH2DialSharingContext{Context: ctx, entered: make(chan struct{}), attemptWait: make(chan struct{})} + results[i], joins[i] = f.reserve(client, witness) + f.wait(witness.entered, "direct reserve follower wait selection") + witnesses = append(witnesses, witness) + } + } + f.joinedWaiters(witnesses) + raw := f.raw(raws) + for i := range cancels { + cancels[i](causes[i]) + } + for i := range results { + f.wait(joins[i], "abandoned reserve caller joined") + if err := (<-results[i]).err; !errors.Is(err, causes[i]) { + t.Errorf("waiter %d error=%v, want its own cause", i, err) + } + } + select { + case <-raw.closed: + t.Error("all caller cancellations closed the client-owned physical initialization before client Close") + default: + } + if err := client.Close(); err != nil { + t.Error(err) + } + select { + case <-raw.closed: + default: + t.Error("client Close returned before actual cold raw Close completed") + } + front.Release() + if accepts, hellos, err := front.state(); accepts != 1 || hellos != 1 || err != nil || entropy.nonceReads.Load() != 0 { + t.Errorf("abandon/Close physical accepts/hellos/error/auth=%d/%d/%v/%d", accepts, hellos, err, entropy.nonceReads.Load()) + } + if _, err := client.DialContext(f.ctx, "tcp", "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) { + t.Errorf("Dial after explicit Close=%v, want net.ErrClosed without target dial", err) + } + t.Log("all direct waiters joined with individual causes; explicit Close and raw-completion checks evaluated") + }) + } +} diff --git a/internal/tunnel/web_h2_initialization_test.go b/internal/tunnel/web_h2_initialization_test.go index 92e0d1d..d9a8201 100644 --- a/internal/tunnel/web_h2_initialization_test.go +++ b/internal/tunnel/web_h2_initialization_test.go @@ -41,9 +41,35 @@ func TestWebH2PostTLSInitializationHonorsCancellation(t *testing.T) { t.Fatal(err) } client.dialer = transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { return wire, nil }) - t.Cleanup(func() { _ = raw.Close(); _ = peer.Close(); _ = client.Close() }) + serverCtx, cancelServer := context.WithTimeout(context.Background(), 3*time.Second) + _ = peer.SetDeadline(time.Now().Add(3 * time.Second)) serverDone := make(chan error, 1) - go func() { serverDone <- tls.Server(peer, serverTLS).HandshakeContext(t.Context()) }() + serverJoined := make(chan struct{}) + dialJoined := make(chan struct{}) + closeDone := make(chan error, 1) + closeJoined := make(chan struct{}) + var closeOnce sync.Once + startClose := func() { + closeOnce.Do(func() { + go func() { + defer close(closeJoined) + closeDone <- client.Close() + }() + }) + } + t.Cleanup(func() { + cancelServer() + _ = raw.Close() + _ = peer.Close() + startClose() + waitWebH2InitializationWorker(t, serverJoined, "TLS peer") + waitWebH2InitializationWorker(t, dialJoined, "caller dial") + waitWebH2InitializationWorker(t, closeJoined, "client Close") + }) + go func() { + defer close(serverJoined) + serverDone <- tls.Server(peer, serverTLS).HandshakeContext(serverCtx) + }() ctx, cancel := context.WithCancelCause(context.Background()) defer cancel(context.Canceled) var dialCtx context.Context = ctx @@ -54,6 +80,17 @@ func TestWebH2PostTLSInitializationHonorsCancellation(t *testing.T) { } dialDone := make(chan error, 1) go func() { + defer close(dialJoined) + if mode == "initialization_deadline" { + // Avoid the public CONNECT wait timer winning the same + // budget first: exercise the physical initializer itself. + session, err := client.openSession(dialCtx) + if session != nil { + _ = closeWebH2Session(session) + } + dialDone <- err + return + } conn, err := client.DialContext(dialCtx, "tcp", "target.invalid:443") if conn != nil { _ = conn.Close() @@ -67,8 +104,13 @@ func TestWebH2PostTLSInitializationHonorsCancellation(t *testing.T) { case <-time.After(time.Second): t.Fatal("dial did not reach the post-TLS initialization write") } - if err := <-serverDone; err != nil { - t.Fatalf("TLS handshake: %v", err) + select { + case err := <-serverDone: + if err != nil { + t.Fatalf("TLS handshake: %v", err) + } + case <-time.After(time.Second): + t.Fatal("TLS peer did not report handshake completion") } client.mu.Lock() registered := len(client.sessions) @@ -85,10 +127,8 @@ func TestWebH2PostTLSInitializationHonorsCancellation(t *testing.T) { want = context.DeadlineExceeded <-dialCtx.Done() case "client_close": - want = context.Canceled - if err := client.Close(); err != nil { - t.Fatal(err) - } + want = net.ErrClosed + startClose() case "initialization_deadline": want = context.DeadlineExceeded } @@ -98,14 +138,24 @@ func TestWebH2PostTLSInitializationHonorsCancellation(t *testing.T) { t.Fatalf("initialization error=%v, want cause %v", err, want) } case <-time.After(time.Second): - _ = raw.Close() - <-dialDone t.Fatal("post-TLS initialization ignored cancellation") } + // Caller cancellation ends only its wait. The physical setup + // remains client-owned, so explicitly Close and join it before + // checking raw-socket and session cleanup in every mode. + startClose() + select { + case err := <-closeDone: + if err != nil { + t.Fatalf("client Close: %v", err) + } + case <-time.After(2 * time.Second): + t.Fatal("client Close did not join physical initialization") + } select { case <-wire.closed: default: - t.Fatal("canceled initializer left its raw connection open") + t.Fatal("joined initialization cleanup left its raw connection open") } client.mu.Lock() registered = len(client.sessions) @@ -172,22 +222,59 @@ func TestWebH2CloseAtInitializationHandoff(t *testing.T) { t.Fatal(err) } client.dialer = transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { return wire, nil }) - t.Cleanup(func() { _ = raw.Close(); _ = peer.Close(); _ = client.Close() }) // Close after the first H2 write fully reaches the peer, but - // before its initializer can publish a session. Either the - // watcher or the closed-client registration gate must win. - wire.onInitialized = func() { _ = client.Close() } - serverDone := make(chan struct{}) + // before Write returns to the initializer. Close must run on + // an independent worker: calling it inside Write would join + // the initializer from that same initializer and deadlock. + handoff := make(chan struct{}) + releaseHandoff := make(chan struct{}) + var releaseOnce sync.Once + release := func() { releaseOnce.Do(func() { close(releaseHandoff) }) } + wire.onInitialized = func() { + close(handoff) + <-releaseHandoff + } + serverCtx, cancelServer := context.WithTimeout(context.Background(), 3*time.Second) + _ = peer.SetDeadline(time.Now().Add(3 * time.Second)) + serverDone := make(chan error, 1) + serverJoined := make(chan struct{}) + dialJoined := make(chan struct{}) + closeDone := make(chan error, 1) + closeJoined := make(chan struct{}) + var closeOnce sync.Once + startClose := func() { + closeOnce.Do(func() { + go func() { + defer close(closeJoined) + closeDone <- client.Close() + }() + }) + } + t.Cleanup(func() { + // Independent socket closure and gate release also work on + // Fatal paths; cleanup never depends on a passing oracle. + release() + cancelServer() + _ = raw.Close() + _ = peer.Close() + startClose() + waitWebH2InitializationWorker(t, serverJoined, "handoff TLS peer") + waitWebH2InitializationWorker(t, dialJoined, "handoff caller dial") + waitWebH2InitializationWorker(t, closeJoined, "handoff client Close") + }) go func() { - defer close(serverDone) + defer close(serverJoined) server := tls.Server(peer, serverTLS) - if server.HandshakeContext(t.Context()) == nil { + err := server.HandshakeContext(serverCtx) + if err == nil { var data [1]byte - _, _ = server.Read(data[:]) + _, err = server.Read(data[:]) } + serverDone <- err }() dialDone := make(chan error, 1) go func() { + defer close(dialJoined) conn, err := client.DialContext(context.Background(), "tcp", "target.invalid:443") if conn != nil { _ = conn.Close() @@ -195,27 +282,62 @@ func TestWebH2CloseAtInitializationHandoff(t *testing.T) { dialDone <- err }() select { + case <-handoff: case err := <-dialDone: - if err == nil { - t.Fatal("closed client published a cold initialized session") - } + t.Fatalf("dial finished before the post-wire handoff: %v", err) case <-time.After(2 * time.Second): - _ = raw.Close() - <-dialDone - t.Fatal("close during initialization handoff did not finish") + t.Fatal("initializer did not reach the post-wire handoff") } - <-serverDone + select { + case err := <-serverDone: + if err != nil { + t.Fatalf("TLS peer did not receive the first H2 bytes: %v", err) + } + case <-time.After(time.Second): + t.Fatal("handoff TLS peer did not report") + } + startClose() select { case <-wire.closed: + case <-time.After(time.Second): + t.Fatal("Close did not cancel the gated initializer's raw socket") + } + select { + case err := <-closeDone: + t.Fatalf("Close returned before initializer handoff was released: %v", err) default: - t.Fatal("closed initializer left its raw connection open") } client.mu.Lock() - registered, current, closed := len(client.sessions), client.current, client.closed + closed := client.closed client.mu.Unlock() if !closed { t.Fatal("test did not reach Close after the initial H2 write") } + release() + select { + case err := <-closeDone: + if err != nil { + t.Fatalf("joined client Close: %v", err) + } + case <-time.After(2 * time.Second): + t.Fatal("Close did not join the released initializer") + } + select { + case err := <-dialDone: + if !errors.Is(err, net.ErrClosed) { + t.Fatalf("closed initializer dial error=%v, want net.ErrClosed", err) + } + case <-time.After(2 * time.Second): + t.Fatal("caller dial did not finish after initialization handoff Close") + } + select { + case <-wire.closed: + default: + t.Fatal("closed initializer left its raw connection open") + } + client.mu.Lock() + registered, current := len(client.sessions), client.current + client.mu.Unlock() if registered != 0 || current != nil { t.Fatal("closed client retained a session after initialization handoff") } @@ -225,6 +347,15 @@ func TestWebH2CloseAtInitializationHandoff(t *testing.T) { } } +func waitWebH2InitializationWorker(t *testing.T, done <-chan struct{}, name string) { + t.Helper() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Errorf("%s worker did not join after independent cleanup", name) + } +} + // The test peer performs a normal verifying TLS 1.3 handshake, then stops // reading. With no client certificate, the first encrypted client flight after // certificate verification contains Finished. Its next write is H2 setup. From d54d334d42196e04389643d071647d4886bf15e7 Mon Sep 17 00:00:00 2001 From: cppla Date: Sat, 3 Oct 2026 08:20:12 +0800 Subject: [PATCH 2/2] fix(web): retain H2 failures for queued callers --- docs/WEB_COVER.md | 3 + internal/tunnel/web_h2_client.go | 82 +++++++--- internal/tunnel/web_h2_dial_queue_test.go | 181 ++++++++++++++++++++++ 3 files changed, 246 insertions(+), 20 deletions(-) create mode 100644 internal/tunnel/web_h2_dial_queue_test.go diff --git a/docs/WEB_COVER.md b/docs/WEB_COVER.md index 05b26fd..b712672 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -384,6 +384,9 @@ Source builds after v1.0.1 share the result of a cold H2 or H3 physical connecti attempt with all callers already waiting on it. A failed handshake does not make those callers start replacement handshakes one after another. A later invocation may retry, subject to the existing `web-auto` cooldown policy. +H2 callers register before queuing for session selection, so they retain +that same completed failure even if it arrives before their selection turn; +new callers after publication can retry without an indefinite failure cache. The shared attempt belongs to the client and retains its configured physical dial timeout; canceling one caller stops only that caller's wait, not the attempt needed by other callers. H2 keeps its TCP-connect timeout separate diff --git a/internal/tunnel/web_h2_client.go b/internal/tunnel/web_h2_client.go index c17c13e..c5dfca5 100644 --- a/internal/tunnel/web_h2_client.go +++ b/internal/tunnel/web_h2_client.go @@ -65,10 +65,11 @@ type WebH2Client struct { mu sync.Mutex closed bool // selected becomes true only after an authenticated CONNECT succeeds. - selected bool - current *webH2ClientSession - sessions map[*webH2ClientSession]struct{} - dial *webH2DialAttempt + selected bool + current *webH2ClientSession + sessions map[*webH2ClientSession]struct{} + dial *webH2DialAttempt + selection *webH2SelectionCohort // Additions are made under mu before the closed gate. Close joins both // unpublished physical setup and cleanup detached from sessions. workers sync.WaitGroup @@ -83,6 +84,14 @@ type webH2DialAttempt struct { err error } +// Register before queuing at dialGate, so a fast completed failure remains +// visible to already queued callers. The worker detaches this cohort when it +// publishes its result; later callers can retry without a failure cache. +// Both the current pointer and its attempt are protected by the client mu. +type webH2SelectionCohort struct { + attempt *webH2DialAttempt +} + type webH2ClientSession struct { raw net.Conn conn webH2TLSClientConn @@ -406,6 +415,21 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati return webH2SessionReservation{}, net.ErrClosed } for { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + return webH2SessionReservation{}, net.ErrClosed + } + if err := contextError(ctx); err != nil { + c.mu.Unlock() + return webH2SessionReservation{}, err + } + cohort := c.selection + if cohort == nil { + cohort = &webH2SelectionCohort{} + c.selection = cohort + } + c.mu.Unlock() select { case c.dialGate <- struct{}{}: case <-ctx.Done(): @@ -436,6 +460,14 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati <-c.dialGate return webH2SessionReservation{}, err } + if attempt := cohort.attempt; attempt != nil { + c.mu.Unlock() + <-c.dialGate + if err := c.waitSessionDial(ctx, attempt); err != nil { + return webH2SessionReservation{}, err + } + continue + } if session := c.current; session != nil { switch session.authState { case webH2ClientAuthFresh: @@ -483,31 +515,38 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati attempt = &webH2DialAttempt{done: make(chan struct{})} c.dial = attempt c.workers.Add(1) + cohort.attempt = attempt go c.runSessionDial(attempt) + } else { + cohort.attempt = attempt } c.mu.Unlock() <-c.dialGate c.closeRetiredSessions(retired) - select { - case <-attempt.done: - case <-c.ctx.Done(): - return webH2SessionReservation{}, net.ErrClosed - case <-ctx.Done(): - if c.ctx.Err() != nil { - return webH2SessionReservation{}, net.ErrClosed - } - return webH2SessionReservation{}, context.Cause(ctx) - } - if c.ctx.Err() != nil { - return webH2SessionReservation{}, net.ErrClosed - } - if err := contextError(ctx); err != nil { + if err := c.waitSessionDial(ctx, attempt); err != nil { return webH2SessionReservation{}, err } - if attempt.err != nil { - return webH2SessionReservation{}, attempt.err + } +} + +func (c *WebH2Client) waitSessionDial(ctx context.Context, attempt *webH2DialAttempt) error { + select { + case <-attempt.done: + case <-c.ctx.Done(): + return net.ErrClosed + case <-ctx.Done(): + if c.ctx.Err() != nil { + return net.ErrClosed } + return context.Cause(ctx) } + if c.ctx.Err() != nil { + return net.ErrClosed + } + if err := contextError(ctx); err != nil { + return err + } + return attempt.err } func (c *WebH2Client) runSessionDial(attempt *webH2DialAttempt) { @@ -534,6 +573,9 @@ func (c *WebH2Client) runSessionDial(attempt *webH2DialAttempt) { c.mu.Lock() attempt.err = err c.dial = nil + if c.selection != nil && c.selection.attempt == attempt { + c.selection = nil + } close(attempt.done) c.mu.Unlock() } diff --git a/internal/tunnel/web_h2_dial_queue_test.go b/internal/tunnel/web_h2_dial_queue_test.go new file mode 100644 index 0000000..8e838ab --- /dev/null +++ b/internal/tunnel/web_h2_dial_queue_test.go @@ -0,0 +1,181 @@ +package tunnel + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +func TestWebH2DialQueuedCallersShareCompletedFailure(t *testing.T) { + f := newWebH2DialSharingFixture(t) + serverTLS, clientTLS := testTLSConfigs(t) + target := f.target() + var destinationDials, covers, unexpected atomic.Int32 + server, err := ListenWebH2(WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: webTestToken, TLSConfig: serverTLS, + Cover: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { covers.Add(1); http.NotFound(w, r) }), + Dialer: transport.DialFunc(func(ctx context.Context, network, address string) (net.Conn, error) { + destinationDials.Add(1) + if network != "tcp" || address != target { + return nil, errors.New("queue fixture refuses a non-owned destination") + } + return (&net.Dialer{Timeout: time.Second}).DialContext(ctx, "tcp4", target) + }), + }) + if err != nil { + t.Fatal(err) + } + type observedRequest struct { + state *webServerConnectionAuth + bearer string + } + requests := make(chan observedRequest, 2) + original := server.server.Handler + server.server.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + joined := make(chan struct{}) + f.register(joined) + defer close(joined) + state, _ := r.Context().Value(webServerConnectionAuthContextKey{}).(*webServerConnectionAuth) + select { + case requests <- observedRequest{state: state, bearer: r.Header.Get("Proxy-Authorization")}: + default: + unexpected.Add(1) + } + original.ServeHTTP(w, r) + }) + serveJoined := make(chan struct{}) + var serveErr error + f.register(serveJoined) + f.checks = append(f.checks, func() { + if serveErr != nil { + t.Errorf("owned queue-control H2 Serve: %v", serveErr) + } + }) + f.close = append(f.close, func() { + if err := server.Close(); err != nil { + t.Errorf("owned queue-control H2 server Close: %v", err) + } + }) + go func() { defer close(serveJoined); serveErr = server.Serve(f.ctx) }() + client, entropy, _ := f.client(FingerprintNative, server.Addr().String(), clientTLS) + realDialer := client.dialer + entered, releaseFailure := make(chan struct{}), make(chan struct{}) + markEntered := sync.OnceFunc(func() { close(entered) }) + finishFailure := sync.OnceFunc(func() { close(releaseFailure) }) + f.close = append(f.close, finishFailure) + fault := errors.New("private controlled first physical failure") + var healthy atomic.Bool + var physicalDials atomic.Int32 + client.dialer = transport.DialFunc(func(ctx context.Context, network, address string) (net.Conn, error) { + physicalDials.Add(1) + if network != "tcp" || address != server.Addr().String() { + return nil, errors.New("queue fixture refuses a non-owned physical endpoint") + } + if healthy.Load() { + return realDialer.DialContext(ctx, network, address) + } + markEntered() + select { + case <-releaseFailure: + case <-ctx.Done(): + } + // Explicit synthetic fast failure, not a TCP performance observation. + return nil, fault + }) + leader, leaderJoined := f.reserve(client, f.ctx) + f.wait(entered, "actual first physical DialFunc entry") + // The existing selection gate is an explicit scheduling barrier. The + // first attempt is already running, but these callers cannot be admitted + // until after that attempt has actually published its completed result. + select { + case client.dialGate <- struct{}{}: + case <-f.ctx.Done(): + t.Fatal("first reserve did not release the selection gate") + } + releaseGate := sync.OnceFunc(func() { <-client.dialGate }) + f.close = append(f.close, releaseGate) // Fatal cleanup always releases it. + queued := make([]<-chan webH2DialSharingResult, 2) + joins := make([]<-chan struct{}, 2) + for i := range queued { + ctx := &webH2DialSharingContext{Context: f.ctx, entered: make(chan struct{}), attemptWait: make(chan struct{})} + queued[i], joins[i] = f.reserve(client, ctx) + f.wait(ctx.entered, "actual queued caller's gate-wait Done evaluation") + } + finishFailure() + f.wait(leaderJoined, "first physical failure published to initiating caller") + firstError := (<-leader).err + if !errors.Is(firstError, fault) { + t.Fatalf("first physical result=%v, want its exact private cause", firstError) + } + client.mu.Lock() + cleared := client.dial == nil && client.current == nil && len(client.sessions) == 0 + client.mu.Unlock() + if !cleared || physicalDials.Load() != 1 || len(client.dialGate) != 1 { + t.Fatalf("completed failure while gate held: cleared=%t physical=%d gate=%d", cleared, physicalDials.Load(), len(client.dialGate)) + } + releaseGate() + for i := range queued { + f.wait(joins[i], "previously queued reserve caller joined") + if err := (<-queued[i]).err; err != firstError { + t.Errorf("queued caller %d missed the completed immutable failure: got=%v want same object=%v", i, err, firstError) + } + } + failedDials := physicalDials.Load() + if failedDials != 1 || entropy.nonceReads.Load() != 0 || destinationDials.Load() != 0 { + t.Errorf("queued failure made physical/auth/destination attempts=%d/%d/%d, want 1/0/0", failedDials, entropy.nonceReads.Load(), destinationDials.Load()) + } + + // A genuinely later invocation may retry; a failure is not cached forever. + // Both flows use real verifying TLS/H2 and the same owned TCP echo target. + healthy.Store(true) + var physical *webH2ClientSession + for round := range 2 { + ctx, cancel := context.WithTimeout(f.ctx, time.Second) + conn, err := client.DialContext(ctx, "tcp", target) + cancel() + if err != nil { + t.Fatalf("later healthy public CONNECT %d: %v", round, err) + } + if err := conn.SetDeadline(time.Now().Add(time.Second)); err != nil { + _ = conn.Close() + t.Fatal(err) + } + payload := []byte(fmt.Sprintf("owned-queued-failure-retry-%d", round)) + _, writeErr := conn.Write(payload) + got := make([]byte, len(payload)) + _, readErr := io.ReadFull(conn, got) + closeErr := conn.Close() + if writeErr != nil || readErr != nil || closeErr != nil || string(got) != string(payload) { + t.Fatalf("later healthy echo %d=%q, write/read/close=%v/%v/%v", round, got, writeErr, readErr, closeErr) + } + client.mu.Lock() + current := client.current + client.mu.Unlock() + if current == nil || (round == 1 && current != physical) { + t.Fatal("later healthy flows did not reuse one physical session") + } + physical = current + } + var observed [2]observedRequest + for i := range observed { + select { + case observed[i] = <-requests: + case <-f.ctx.Done(): + t.Fatal("healthy handler observation exceeded fixture budget") + } + } + _, short := parseWebSessionBearer(observed[1].bearer) + if observed[0].state == nil || observed[0].state != observed[1].state || len(observed[0].bearer) < 385 || !short || entropy.nonceReads.Load() != 1 || physicalDials.Load() != failedDials+1 || destinationDials.Load() != 2 || covers.Load() != 0 || unexpected.Load() != 0 { + t.Errorf("healthy retry/full-short auth violated: lengths=%d/%d short=%t physical/nonce/destination/cover/unexpected=%d/%d/%d/%d/%d", len(observed[0].bearer), len(observed[1].bearer), short, physicalDials.Load(), entropy.nonceReads.Load(), destinationDials.Load(), covers.Load(), unexpected.Load()) + } + t.Log("queued-before-completion failure and independent real TLS/H2 full/short-auth echo controls evaluated") +}