From 949761b8045714863fdacc6b747b58318f340abd Mon Sep 17 00:00:00 2001 From: cppla Date: Tue, 22 Sep 2026 20:32:34 +0800 Subject: [PATCH] fix: preserve active tunnels and bound connection retirement --- docs/DEPLOYMENT.md | 6 + docs/WEB_COVER.md | 6 + integration/proxy_activity_test.go | 211 +++++++++++++++ internal/proxy/config.go | 5 +- internal/proxy/http.go | 5 +- internal/proxy/relay.go | 72 +++-- internal/proxy/relay_test.go | 109 ++++++++ internal/tunnel/web_client.go | 20 +- internal/tunnel/web_client_test.go | 6 + internal/tunnel/web_client_udp_health_test.go | 182 +++++++++++++ internal/tunnel/web_h2_client.go | 77 ++++-- internal/tunnel/web_h2_lifecycle_test.go | 249 ++++++++++++++++++ internal/tunnel/web_stream.go | 11 +- internal/tunnel/web_udp.go | 25 +- 14 files changed, 931 insertions(+), 53 deletions(-) create mode 100644 integration/proxy_activity_test.go create mode 100644 internal/tunnel/web_client_udp_health_test.go create mode 100644 internal/tunnel/web_h2_lifecycle_test.go diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index e223dbf..fbce9ec 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -244,6 +244,12 @@ behavior and limits are in [WEB_COVER.md](WEB_COVER.md). Default local endpoints are loopback-only SOCKS5 `127.0.0.1:1080` and HTTP `127.0.0.1:8080`. Use `socks5h://` when the relay should resolve names. +For SOCKS5 TCP and HTTP CONNECT tunnels, `--idle-timeout` (default `5m`) +measures inactivity across both directions: an active download or upload does +not need reverse-direction application traffic to stay open. A blocked write +still has its own timeout, so an unresponsive receiver cannot retain a tunnel +indefinitely. Ordinary forwarded HTTP bodies keep per-operation timeouts. + Before leaving a client running, verify a real authenticated relay path with the same connection flags: diff --git a/docs/WEB_COVER.md b/docs/WEB_COVER.md index a97982a..08f46e0 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -165,6 +165,12 @@ both pacing fields as `not-applicable`, and `--pacing=fixed-rate` is rejected. `h2` does not advertise UDP. The H3-to-H2 fallback carries only new TCP streams: UDP never falls back to H2 and fails when H3 is unavailable. +Only a newly authenticated CONNECT-UDP response is fresh evidence that the H3 +path has recovered. Sending on a cached UDP target merely queues a datagram +locally; it does not clear the TCP fallback cooldown or change the last +successful transport. A newly authenticated target rejection proves path +health without being counted as a successful target connection. + ## CONNECT-UDP boundary The H3 UDP path follows diff --git a/integration/proxy_activity_test.go b/integration/proxy_activity_test.go new file mode 100644 index 0000000..152e7fd --- /dev/null +++ b/integration/proxy_activity_test.go @@ -0,0 +1,211 @@ +package integration_test + +import ( + "bufio" + "bytes" + "context" + "fmt" + "io" + "net" + "net/http" + "strings" + "testing" + "time" + + "github.com/cppla/autocar/internal/proxy" + "github.com/cppla/autocar/internal/transport" + "github.com/cppla/autocar/internal/tunnel" +) + +// A completed upload must not leave a write-deadline timer armed on the H2 +// stream, and its quiet read direction must not expire an active download. +// Unlike an echo or a half-close test, both directions stay open throughout. +func TestHTTPConnectWebH2ActiveDownloadSurvivesIdleTimeout(t *testing.T) { + const idleTimeout = 250 * time.Millisecond + const interval = 50 * time.Millisecond + request := []byte("go!") + chunk := []byte("download-chunk\x00\xff") + payload := bytes.Repeat(chunk, 16) + dialer := startActivityWebH2Client(t) + + for _, pipelined := range []bool{false, true} { + name := "request_after_connect" + if pipelined { + name = "request_with_connect" + } + t.Run(name, func(t *testing.T) { + target, targetResult := startActivityDownloadTarget(t, request, chunk, 16, interval) + proxyAddress := startActivityHTTPProxy(t, dialer, idleTimeout) + conn := dialWithDeadline(t, proxyAddress) + defer conn.Close() + head := fmt.Sprintf("CONNECT %s HTTP/1.1\r\nHost: %s\r\n\r\n", target, target) + initial := []byte(head) + if pipelined { + initial = append(initial, request...) + } + started := time.Now() + writeFull(t, conn, initial) + reader := bufio.NewReader(conn) + status, err := reader.ReadString('\n') + if err != nil || !strings.Contains(status, " 200 ") { + t.Fatalf("CONNECT response = %q, err = %v", status, err) + } + for { + line, err := reader.ReadString('\n') + if err != nil { + t.Fatalf("read CONNECT headers: %v", err) + } + if line == "\r\n" { + break + } + } + if !pipelined { + writeFull(t, conn, request) + } + + got := make([]byte, len(payload)) + n, err := io.ReadFull(reader, got) + if err != nil { + t.Fatalf("active H2 download stopped after %d/%d bytes and %s: %v", n, len(got), time.Since(started), err) + } + if !bytes.Equal(got, payload) { + t.Fatal("download payload differs from target bytes") + } + if elapsed := time.Since(started); elapsed <= 2*idleTimeout { + t.Fatalf("download lasted %s, want more than two idle intervals", elapsed) + } + // The upload side must still work after the long download, without + // relying on EOF or CloseWrite to complete either direction. + writeFull(t, conn, []byte("!")) + select { + case err := <-targetResult: + if err != nil { + t.Fatalf("target exchange: %v", err) + } + case <-time.After(operationTimeout): + t.Fatal("target did not receive final acknowledgment") + } + }) + } +} + +func startActivityWebH2Client(t *testing.T) *tunnel.WebH2Client { + t.Helper() + material := newTLSMaterial(t) + server, err := tunnel.ListenWebH2(tunnel.WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: material.server, + Cover: http.NotFoundHandler(), Dialer: &net.Dialer{Timeout: operationTimeout}, + HandshakeTimeout: operationTimeout, DialTimeout: operationTimeout, + }) + if err != nil { + t.Fatalf("listen activity-test H2 relay: %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- server.Serve(ctx) }() + t.Cleanup(func() { + cancel() + _ = server.Close() + waitForServe(t, "activity-test H2 relay", done) + }) + client, err := tunnel.NewWebH2Client(tunnel.WebH2ClientConfig{ + ServerAddress: server.Addr().String(), Token: testToken, TLSConfig: material.client, + HandshakeTimeout: operationTimeout, DialTimeout: operationTimeout, + }) + if err != nil { + t.Fatalf("create activity-test H2 client: %v", err) + } + t.Cleanup(func() { _ = client.Close() }) + return client +} + +func startActivityHTTPProxy(t *testing.T, dialer transport.Dialer, idleTimeout time.Duration) string { + t.Helper() + server, err := proxy.NewHTTPServer(proxy.Config{ + Dialer: dialer, HandshakeTimeout: operationTimeout, DialTimeout: operationTimeout, + IdleTimeout: idleTimeout, MaxConnections: 4, + }) + if err != nil { + t.Fatalf("create activity-test HTTP proxy: %v", err) + } + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen activity-test HTTP proxy: %v", err) + } + done := make(chan error, 1) + go func() { done <- server.Serve(listener) }() + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), operationTimeout) + defer cancel() + if err := server.Shutdown(ctx); err != nil && !isExpectedClose(err) { + t.Errorf("shut down activity-test HTTP proxy: %v", err) + } + waitForServe(t, "activity-test HTTP proxy", done) + }) + return listener.Addr().String() +} + +func startActivityDownloadTarget(t *testing.T, request, chunk []byte, count int, interval time.Duration) (string, <-chan error) { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen activity-test target: %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + stopped := make(chan struct{}) + go func() { + defer close(stopped) + result <- func() error { + conn, err := listener.Accept() + if err != nil { + return err + } + defer conn.Close() + stop := context.AfterFunc(ctx, func() { _ = conn.Close() }) + defer stop() + if err := conn.SetDeadline(time.Now().Add(operationTimeout)); err != nil { + return err + } + got := make([]byte, len(request)) + if _, err := io.ReadFull(conn, got); err != nil { + return fmt.Errorf("read initial request: %w", err) + } + if !bytes.Equal(got, request) { + return fmt.Errorf("initial request = %q, want %q", got, request) + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + for range count { + select { + case <-ticker.C: + case <-ctx.Done(): + return ctx.Err() + } + if n, err := conn.Write(chunk); err != nil { + return fmt.Errorf("write download chunk: %w", err) + } else if n != len(chunk) { + return io.ErrShortWrite + } + } + ack := make([]byte, 1) + if _, err := io.ReadFull(conn, ack); err != nil { + return fmt.Errorf("read final acknowledgment: %w", err) + } + if ack[0] != '!' { + return fmt.Errorf("final acknowledgment = %q", ack) + } + return nil + }() + }() + t.Cleanup(func() { + cancel() + _ = listener.Close() + select { + case <-stopped: + case <-time.After(operationTimeout): + t.Error("activity-test target did not stop") + } + }) + return listener.Addr().String(), result +} diff --git a/internal/proxy/config.go b/internal/proxy/config.go index 5509f44..d432942 100644 --- a/internal/proxy/config.go +++ b/internal/proxy/config.go @@ -29,7 +29,10 @@ type Config struct { Authenticator Authenticator HandshakeTimeout time.Duration DialTimeout time.Duration - IdleTimeout time.Duration + // IdleTimeout bounds inactivity across either direction of a CONNECT + // tunnel, and separately bounds a blocked write. Ordinary HTTP request + // and response bodies retain their per-operation inactivity bound. + IdleTimeout time.Duration // MaxConnections is the number of accepted client TCP connections that may // be active at once. Zero uses a conservative default; a negative value is // rejected. diff --git a/internal/proxy/http.go b/internal/proxy/http.go index 4b339a7..a8e360f 100644 --- a/internal/proxy/http.go +++ b/internal/proxy/http.go @@ -329,10 +329,7 @@ func writeGatewayError(w http.ResponseWriter, err error) { func writeAll(conn net.Conn, payload []byte, timeout time.Duration) error { for len(payload) > 0 { - if timeout > 0 { - _ = conn.SetWriteDeadline(time.Now().Add(timeout)) - } - written, err := conn.Write(payload) + written, err := writeWithStallDeadline(conn, payload, timeout) if written < 0 || written > len(payload) { return errors.New("proxy: invalid write count") } diff --git a/internal/proxy/relay.go b/internal/proxy/relay.go index 5aa7e19..8f09746 100644 --- a/internal/proxy/relay.go +++ b/internal/proxy/relay.go @@ -31,12 +31,25 @@ func (c *activityConn) Read(p []byte) (int, error) { } func (c *activityConn) Write(p []byte) (int, error) { - if c.timeout > 0 { - if err := c.Conn.SetWriteDeadline(time.Now().Add(c.timeout)); err != nil { - return 0, err - } + return writeWithStallDeadline(c.Conn, p, c.timeout) +} + +// Bound only the pending write. In particular, H2 streams implement deadlines +// by aborting the stream: leaving a completed write's timer armed would later +// terminate an otherwise active download with no more uploads. +func writeWithStallDeadline(conn net.Conn, p []byte, timeout time.Duration) (int, error) { + if timeout <= 0 { + return conn.Write(p) } - return c.Conn.Write(p) + if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil { + return 0, err + } + n, err := conn.Write(p) + clearErr := conn.SetWriteDeadline(time.Time{}) + if err == nil { + err = clearErr + } + return n, err } func (c *activityConn) CloseWrite() error { @@ -61,11 +74,17 @@ func relay(left, right net.Conn, idleTimeout time.Duration) error { err error } results := make(chan result, 2) + activity := &relayActivity{left: left, right: right, timeout: idleTimeout} + if err := activity.refresh(); err != nil { + _ = left.Close() + _ = right.Close() + return err + } go func() { - results <- result{destination: right, err: copyHalf(right, left, idleTimeout)} + results <- result{destination: right, err: copyHalf(right, left, activity)} }() go func() { - results <- result{destination: left, err: copyHalf(left, right, idleTimeout)} + results <- result{destination: left, err: copyHalf(left, right, activity)} }() var relayErrors []error @@ -101,25 +120,43 @@ func relay(left, right net.Conn, idleTimeout time.Duration) error { return errors.Join(relayErrors...) } -func copyHalf(dst, src net.Conn, idleTimeout time.Duration) error { +// relayActivity makes read inactivity a property of the whole tunnel, not +// either direction independently. Downloads, uploads and server-push streams +// may legitimately have no reverse-direction application data for minutes. +// Serializing the update prevents an older activity event from overwriting a +// newer deadline. Writes retain their own stall deadline for backpressure. +type relayActivity struct { + mu sync.Mutex + left, right net.Conn + timeout time.Duration +} + +func (a *relayActivity) refresh() error { + if a.timeout <= 0 { + return nil + } + a.mu.Lock() + defer a.mu.Unlock() + deadline := time.Now().Add(a.timeout) + return errors.Join(a.left.SetReadDeadline(deadline), a.right.SetReadDeadline(deadline)) +} + +func copyHalf(dst, src net.Conn, activity *relayActivity) error { bufferPtr := relayBufferPool.Get().(*[]byte) defer relayBufferPool.Put(bufferPtr) buffer := *bufferPtr for { - if idleTimeout > 0 { - _ = src.SetReadDeadline(time.Now().Add(idleTimeout)) - } n, readErr := src.Read(buffer) if n < 0 || n > len(buffer) { return errors.New("proxy: invalid read count") } if n > 0 { + if err := activity.refresh(); err != nil { + return err + } written := 0 for written < n { - if idleTimeout > 0 { - _ = dst.SetWriteDeadline(time.Now().Add(idleTimeout)) - } - m, writeErr := dst.Write(buffer[written:n]) + m, writeErr := writeWithStallDeadline(dst, buffer[written:n], activity.timeout) if m < 0 || m > n-written { return errors.New("proxy: invalid write count") } @@ -127,6 +164,11 @@ func copyHalf(dst, src net.Conn, idleTimeout time.Duration) error { if writeErr != nil { return writeErr } + if m > 0 { + if err := activity.refresh(); err != nil { + return err + } + } if m == 0 { return io.ErrShortWrite } diff --git a/internal/proxy/relay_test.go b/internal/proxy/relay_test.go index c4a3320..b981657 100644 --- a/internal/proxy/relay_test.go +++ b/internal/proxy/relay_test.go @@ -76,6 +76,115 @@ func TestRelayCleanEOFPreservesReverseDirection(t *testing.T) { } } +func TestRelayOneWayActivityKeepsConnectionAlive(t *testing.T) { + for _, direction := range []string{"upload", "download"} { + t.Run(direction, func(t *testing.T) { + leftRelay, leftPeer := net.Pipe() + rightRelay, rightPeer := net.Pipe() + defer leftRelay.Close() + defer rightRelay.Close() + defer leftPeer.Close() + defer rightPeer.Close() + const idle = 200 * time.Millisecond + done := make(chan error, 1) + go func() { done <- relay(leftRelay, rightRelay, idle) }() + sender, receiver := leftPeer, rightPeer + if direction == "download" { + sender, receiver = rightPeer, leftPeer + } + _ = sender.SetDeadline(time.Now().Add(3 * time.Second)) + _ = receiver.SetDeadline(time.Now().Add(3 * time.Second)) + written := make(chan error, 1) + go func() { + for i := 0; i < 16; i++ { + if _, err := sender.Write([]byte{byte(i)}); err != nil { + written <- err + return + } + time.Sleep(idle / 5) + } + written <- nil + }() + for i := 0; i < 16; i++ { + var payload [1]byte + if _, err := io.ReadFull(receiver, payload[:]); err != nil { + t.Fatalf("active %s stopped at byte %d: %v", direction, i, err) + } + if payload[0] != byte(i) { + t.Fatalf("payload = %d, want %d", payload[0], i) + } + } + if err := <-written; err != nil { + t.Fatal(err) + } + // Once both directions stop making progress, the same idle bound + // must still release the connection and its two relay workers. + select { + case err := <-done: + var timeout net.Error + if !errors.As(err, &timeout) || !timeout.Timeout() { + t.Fatalf("idle relay error = %v, want timeout", err) + } + case <-time.After(3 * idle): + t.Fatal("inactive relay did not stop") + } + }) + } +} + +func TestRelayBlockedWriteTimesOutDespiteReverseActivity(t *testing.T) { + leftRelay, leftPeer := net.Pipe() + rightRelay, rightPeer := net.Pipe() + defer leftRelay.Close() + defer rightRelay.Close() + defer leftPeer.Close() + defer rightPeer.Close() + const idle = 200 * time.Millisecond + done := make(chan error, 1) + go func() { done <- relay(leftRelay, rightRelay, idle) }() + _ = leftPeer.SetDeadline(time.Now().Add(2 * time.Second)) + _ = rightPeer.SetDeadline(time.Now().Add(2 * time.Second)) + // No one reads rightPeer: forwarding this upload must remain bounded, + // even while rightPeer keeps supplying successful reverse-direction data. + mustWrite(t, leftPeer, []byte("blocked upload")) + written := make(chan struct{}) + go func() { + defer close(written) + for { + if _, err := rightPeer.Write([]byte("x")); err != nil { + return + } + time.Sleep(idle / 10) + } + }() + var count int + for { + var payload [1]byte + n, err := leftPeer.Read(payload[:]) + count += n + if err != nil { + break + } + } + if count == 0 { + t.Fatal("reverse direction made no progress") + } + select { + case err := <-done: + var timeout net.Error + if !errors.As(err, &timeout) || !timeout.Timeout() { + t.Fatalf("blocked write error = %v, want timeout", err) + } + case <-time.After(3 * idle): + t.Fatal("blocked write was kept alive by unrelated activity") + } + select { + case <-written: + case <-time.After(time.Second): + t.Fatal("reverse writer did not exit") + } +} + type readErrorConn struct { net.Conn err error diff --git a/internal/tunnel/web_client.go b/internal/tunnel/web_client.go index 71493ff..f318503 100644 --- a/internal/tunnel/web_client.go +++ b/internal/tunnel/web_client.go @@ -492,24 +492,20 @@ func (c *WebClient) DialPacket(ctx context.Context) (transport.PacketConn, error _ = packet.Close() return nil, net.ErrClosed } - return &webClientPacketConn{parent: c, inner: packet}, nil + if webPacket, ok := packet.(*webUDPPacketConn); ok { + webPacket.setAuthenticationObserver(c.recordOutOfBandPrimaryHealthy) + } + return &webClientPacketConn{inner: packet}, nil } type webClientPacketConn struct { - parent *WebClient - inner transport.PacketConn + inner transport.PacketConn } func (c *webClientPacketConn) Send(payload []byte, address string) error { - if err := c.inner.Send(payload, address); err != nil { - var connectErr *WebConnectError - if errors.As(err, &connectErr) { - c.parent.recordOutOfBandPrimaryHealthy(false) - } - return err - } - c.parent.recordOutOfBandPrimaryHealthy(true) - return nil + // A cached UDP session only queues this datagram locally. It does not prove + // current path health; only a newly authenticated CONNECT-UDP response does. + return c.inner.Send(payload, address) } func (c *webClientPacketConn) Receive() ([]byte, string, error) { return c.inner.Receive() } diff --git a/internal/tunnel/web_client_test.go b/internal/tunnel/web_client_test.go index 05a3359..6410773 100644 --- a/internal/tunnel/web_client_test.go +++ b/internal/tunnel/web_client_test.go @@ -648,6 +648,9 @@ func TestWebClientAuthenticatedUDPSuccessClosesCircuitAgainstOlderTCPFailure(t * if err := packet.Send([]byte("healthy"), "example.com:53"); err != nil { t.Fatal(err) } + // This fake has no HTTP/3 handshake. Deliver its verified-response event + // explicitly; a successful local Send alone must never deliver that event. + client.recordOutOfBandPrimaryHealthy(true) _ = packet.Close() close(releaseProbe) if err := <-probeDone; err != nil { @@ -745,6 +748,9 @@ func TestWebClientAuthenticatedUDPErrorClosesCircuitAgainstOlderTCPFailure(t *te if !errors.As(err, &gotTargetFailure) || gotTargetFailure != targetFailure { t.Fatalf("packet error = %v, want exact authenticated target error", err) } + // Model a newly verified rejection, not a replayed error from a cached + // packet implementation. Real H3 observer wiring is tested separately. + client.recordOutOfBandPrimaryHealthy(false) _ = packet.Close() if got := client.SelectedTransport(); got != webAuthTransportH2 { t.Fatalf("failed CONNECT-UDP changed selection to %q", got) diff --git a/internal/tunnel/web_client_udp_health_test.go b/internal/tunnel/web_client_udp_health_test.go new file mode 100644 index 0000000..4df29b1 --- /dev/null +++ b/internal/tunnel/web_client_udp_health_test.go @@ -0,0 +1,182 @@ +package tunnel + +import ( + "context" + "errors" + "net" + "net/netip" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +func assertWebPrimaryFailureGeneration(t *testing.T, client *WebClient, generation uint64) { + t.Helper() + client.mu.Lock() + defer client.mu.Unlock() + if client.primaryFailedAt.IsZero() || client.primaryStateID != generation { + t.Fatalf("H3 failure changed without a new authenticated response: failed=%v generation=%d, want %d", + client.primaryFailedAt, client.primaryStateID, generation) + } +} + +func markWebPrimaryFailed(t *testing.T, client *WebClient, now time.Time) uint64 { + t.Helper() + client.mu.Lock() + client.nextPrimaryID++ + generation := client.nextPrimaryID + client.mu.Unlock() + if !client.recordPrimaryFailure(generation, now) { + t.Fatal("could not record H3 failure") + } + return generation +} + +func TestWebClientBufferedUDPSendDoesNotRecoverFailedH3Circuit(t *testing.T) { + clock := newWebClientTestClock(time.Unix(600, 0)) + basePrimary := &webClientTestDialer{dial: func(context.Context, int64, string, string) (net.Conn, error) { + return nil, errors.New("H3 path unreachable") + }} + // Exercise the real cached-target Send path. There is deliberately no + // sender goroutine or peer: successful Send can only enqueue locally. + session := &webUDPClientSession{done: make(chan struct{}), outbound: make(chan []byte, 1)} + inner := &webUDPPacketConn{sessions: map[string]*webUDPClientSession{"example.com:53": session}} + primary := &webClientTestPacketDialer{webClientTestDialer: basePrimary, packet: inner} + fallback := &webClientTestDialer{dial: func(context.Context, int64, string, string) (net.Conn, error) { + return closedWebClientTestConn(), nil + }} + client, err := newWebClientWithPaths( + webClientPath{name: webAuthTransportH3, dialer: primary}, + webClientPath{name: webAuthTransportH2, dialer: fallback}, + time.Second, time.Minute, clock.Now, + ) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + first, err := client.DialContext(context.Background(), "tcp", "example.com:443") + if err != nil { + t.Fatal(err) + } + _ = first.Close() + client.mu.Lock() + generation := client.primaryStateID + client.mu.Unlock() + packet, err := client.DialPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + assertWebPrimaryFailureGeneration(t, client, generation) + if err := packet.Send([]byte("buffered only"), "example.com:53"); err != nil { + t.Fatal(err) + } + if len(session.outbound) != 1 { + t.Fatal("datagram was not locally buffered") + } + assertWebPrimaryFailureGeneration(t, client, generation) + if err := packet.Send([]byte("queue full"), "example.com:53"); !errors.Is(err, transport.ErrPacketQueueFull) { + t.Fatalf("full local queue error = %v", err) + } + assertWebPrimaryFailureGeneration(t, client, generation) + second, err := client.DialContext(context.Background(), "tcp", "example.com:443") + if err != nil { + t.Fatal(err) + } + _ = second.Close() + if got := basePrimary.calls.Load(); got != 1 { + t.Fatalf("local UDP enqueue reopened failed H3 circuit: primary calls=%d, want 1 before cooldown", got) + } + if got := client.SelectedTransport(); got != webAuthTransportH2 { + t.Fatalf("local UDP enqueue changed selected transport to %q", got) + } +} + +func TestWebClientH3UDPOnlyNewAuthenticatedResponseRecoversCircuit(t *testing.T) { + target := startWebUDPEcho(t) + resolver := newWebUDPTestResolver(map[string]netip.AddrPort{target.String(): target}) + h3 := newWebH3RetirementClient(t, 4, nil, resolver) + fallback := &webClientTestDialer{dial: func(context.Context, int64, string, string) (net.Conn, error) { + return closedWebClientTestConn(), nil + }} + clock := newWebClientTestClock(time.Unix(700, 0)) + client, err := newWebClientWithPaths( + webClientPath{name: webAuthTransportH3, dialer: h3}, + webClientPath{name: webAuthTransportH2, dialer: fallback}, + time.Second, time.Minute, clock.Now, + ) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + generation := markWebPrimaryFailed(t, client, clock.Now()) + packet, err := client.DialPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packet.Close() }) + assertWebPrimaryFailureGeneration(t, client, generation) + assertWebUDPEcho(t, packet, []byte("new authenticated CONNECT-UDP"), target.String()) + client.mu.Lock() + recovered := client.primaryFailedAt.IsZero() && client.primaryStateID > generation + client.mu.Unlock() + if !recovered || client.SelectedTransport() != webAuthTransportH3 { + t.Fatal("new authenticated H3 UDP response did not recover the circuit") + } + generation = markWebPrimaryFailed(t, client, clock.Now()) + assertWebUDPEcho(t, packet, []byte("same cached target"), target.String()) + assertWebPrimaryFailureGeneration(t, client, generation) + // A newly authenticated target rejection still proves path health, without + // selecting that failed target as a successful H3 stream. + fallbackConn, err := client.dialFallback(context.Background(), "tcp", "example.com:443", nil) + if err != nil { + t.Fatal(err) + } + _ = fallbackConn.Close() + err = packet.Send([]byte("reject target"), "missing.example:53") + var rejection *WebConnectError + if !errors.As(err, &rejection) { + t.Fatalf("expected authenticated H3 target rejection, got %v", err) + } + client.mu.Lock() + recovered = client.primaryFailedAt.IsZero() && client.primaryStateID > generation + client.mu.Unlock() + if !recovered || client.SelectedTransport() != webAuthTransportH2 { + t.Fatal("authenticated target rejection did not restore health while preserving H2 selection") + } +} + +func TestWebClientH3UDPAuthenticationFailureDoesNotRecoverCircuit(t *testing.T) { + good := newWebH3RetirementClient(t, 4, nil, newWebUDPTestResolver(map[string]netip.AddrPort{})) + h3, err := NewWebH3Client(WebH3ClientConfig{ + ServerAddress: good.address, Token: webTestToken + "-wrong", TLSConfig: good.tlsConfig, + FingerprintProfile: H3FingerprintNative, DialTimeout: time.Second, HandshakeTimeout: time.Second, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = h3.Close() }) + fallback := &webClientTestDialer{dial: func(context.Context, int64, string, string) (net.Conn, error) { + return closedWebClientTestConn(), nil + }} + clock := newWebClientTestClock(time.Unix(800, 0)) + client, err := newWebClientWithPaths( + webClientPath{name: webAuthTransportH3, dialer: h3}, + webClientPath{name: webAuthTransportH2, dialer: fallback}, + time.Second, time.Minute, clock.Now, + ) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + generation := markWebPrimaryFailed(t, client, clock.Now()) + packet, err := client.DialPacket(context.Background()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packet.Close() }) + if err := packet.Send([]byte("wrong credential"), "missing.example:53"); err == nil { + t.Fatal("incorrect credentials were accepted") + } + assertWebPrimaryFailureGeneration(t, client, generation) +} diff --git a/internal/tunnel/web_h2_client.go b/internal/tunnel/web_h2_client.go index 93bce64..5bc7431 100644 --- a/internal/tunnel/web_h2_client.go +++ b/internal/tunnel/web_h2_client.go @@ -58,6 +58,7 @@ type WebH2Client struct { } type webH2ClientSession struct { + raw net.Conn conn webH2TLSClientConn h2 *http2.ClientConn authState webH2ClientAuthState @@ -66,6 +67,9 @@ type webH2ClientSession struct { // opening protects a selected session while an authentication ticket waits // for replay-window capacity, before RoundTrip owns an HTTP/2 stream. opening int + // active counts returned net.Conns until full close, including deadline + // cancellation. Half-closes keep the opposite direction's ownership. + active int } type webH2ClientAuthState uint8 @@ -325,14 +329,17 @@ func (c *WebH2Client) DialContext(ctx context.Context, network, address string) return nil, net.ErrClosed } c.selected = true + session.active++ c.mu.Unlock() - return newWebH2Conn( + conn := newWebH2Conn( response.Body, requestWriter, cancelStream, session.conn.LocalAddr(), session.conn.RemoteAddr(), - ), nil + ) + conn.onClose = func() { c.releaseSessionStream(session) } + return conn, nil } func normalizeWebH2Authority(address string) (string, error) { @@ -384,7 +391,6 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati <-c.dialGate return webH2SessionReservation{}, net.ErrClosed } - c.cleanupIdleSessionsLocked() if session := c.current; session != nil { switch session.authState { case webH2ClientAuthReady: @@ -412,7 +418,9 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati c.current = nil } } + retired := c.cleanupIdleSessionsLocked() c.mu.Unlock() + closeWebH2Sessions(retired) session, err := c.openSession(ctx) if err != nil { @@ -421,7 +429,7 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati } if !session.h2.CanTakeNewRequest() { <-c.dialGate - _ = session.h2.Close() + _ = closeWebH2Session(session) return webH2SessionReservation{}, errors.New("tunnel: new web-cover HTTP/2 connection rejected its first stream") } @@ -429,7 +437,7 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati if c.closed { c.mu.Unlock() <-c.dialGate - _ = session.h2.Close() + _ = closeWebH2Session(session) return webH2SessionReservation{}, net.ErrClosed } c.current = session @@ -444,8 +452,17 @@ func (c *WebH2Client) reserveSession(ctx context.Context) (webH2SessionReservati func (c *WebH2Client) releaseSessionReservation(session *webH2ClientSession) { c.mu.Lock() session.opening-- - c.cleanupIdleSessionsLocked() + retired := c.cleanupIdleSessionsLocked() + c.mu.Unlock() + closeWebH2Sessions(retired) +} + +func (c *WebH2Client) releaseSessionStream(session *webH2ClientSession) { + c.mu.Lock() + session.active-- + retired := c.cleanupIdleSessionsLocked() c.mu.Unlock() + closeWebH2Sessions(retired) } func (c *WebH2Client) openSession(ctx context.Context) (*webH2ClientSession, error) { @@ -487,6 +504,7 @@ func (c *WebH2Client) openSession(ctx context.Context) (*webH2ClientSession, err return nil, fmt.Errorf("tunnel: initialize web-cover HTTP/2 connection: %w", err) } return &webH2ClientSession{ + raw: raw, conn: tlsConn, h2: clientConn, authState: webH2ClientAuthBootstrapping, @@ -528,9 +546,8 @@ func (c *WebH2Client) failSessionAuthentication(session *webH2ClientSession) { c.current = nil } delete(c.sessions, session) - session.auth.close() c.mu.Unlock() - _ = session.h2.Close() + _ = closeWebH2Session(session) } func (c *WebH2Client) noteSessionFailure(session *webH2ClientSession) { @@ -538,23 +555,46 @@ func (c *WebH2Client) noteSessionFailure(session *webH2ClientSession) { if c.current == session && !session.h2.CanTakeNewRequest() { c.current = nil } - c.cleanupIdleSessionsLocked() + retired := c.cleanupIdleSessionsLocked() c.mu.Unlock() + closeWebH2Sessions(retired) } -func (c *WebH2Client) cleanupIdleSessionsLocked() { +// Only inspect ownership maintained under c.mu. http2.ClientConn.State takes +// the HTTP/2 write mutex and can wait indefinitely behind a stalled socket; +// using it here would prevent even client Close from interrupting that socket. +// Closing the detached sessions must also happen outside c.mu. +func (c *WebH2Client) cleanupIdleSessionsLocked() []*webH2ClientSession { + var retired []*webH2ClientSession for session := range c.sessions { - if session == c.current { - continue - } - state := session.h2.State() - if session.opening != 0 || state.StreamsActive != 0 || state.StreamsReserved != 0 || state.StreamsPending != 0 { + if session == c.current || session.opening != 0 || session.active != 0 { continue } - _ = session.h2.Close() - session.auth.close() delete(c.sessions, session) + retired = append(retired, session) + } + return retired +} + +func closeWebH2Sessions(sessions []*webH2ClientSession) { + for _, session := range sessions { + _ = closeWebH2Session(session) + } +} + +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 + } } + return errors.Join(rawErr, session.h2.Close()) } // Close prevents future dials, cancels active streams, and closes every pooled @@ -569,7 +609,6 @@ func (c *WebH2Client) Close() error { c.cancel() sessions := make([]*webH2ClientSession, 0, len(c.sessions)) for session := range c.sessions { - session.auth.close() sessions = append(sessions, session) } c.current = nil @@ -578,7 +617,7 @@ func (c *WebH2Client) Close() error { var result error for _, session := range sessions { - result = errors.Join(result, session.h2.Close()) + result = errors.Join(result, closeWebH2Session(session)) } return result } diff --git a/internal/tunnel/web_h2_lifecycle_test.go b/internal/tunnel/web_h2_lifecycle_test.go new file mode 100644 index 0000000..4d9eefc --- /dev/null +++ b/internal/tunnel/web_h2_lifecycle_test.go @@ -0,0 +1,249 @@ +package tunnel + +import ( + "context" + "errors" + "io" + "net" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" + "golang.org/x/net/http2" +) + +// A peer that has stopped reading leaves a real HTTP/2 SETTINGS ACK blocked +// while holding x/net's write mutex. State() must not be used for retirement: +// it takes that same mutex, unlike CanTakeNewRequest's memory-only state lock. +func TestWebH2RetirementAndCloseDoNotWaitForWireWrite(t *testing.T) { + for _, active := range []int{0, 1} { + t.Run(map[int]string{0: "drained", 1: "active"}[active], func(t *testing.T) { + wire, h2 := newWebH2StalledWriter(t) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + session := &webH2ClientSession{raw: wire, h2: h2, opening: 1, active: active} + client := &WebH2Client{ + ctx: ctx, cancel: cancel, dialGate: make(chan struct{}, 1), + sessions: map[*webH2ClientSession]struct{}{session: {}}, + } + // Leave the dial gate occupied so a second dial can only finish + // through its deadline, without accessing the network. + client.dialGate <- struct{}{} + released := make(chan struct{}) + go func() { client.releaseSessionReservation(session); close(released) }() + select { + case <-released: + case <-time.After(time.Second): + _ = wire.Close() + <-released + t.Fatal("retirement waited for the HTTP/2 write mutex") + } + if got := wire.closes.Load(); (got != 0) != (active == 0) { + t.Fatalf("wire closes=%d with active=%d", got, active) + } + dialCtx, dialCancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer dialCancel() + dialDone := make(chan error, 1) + go func() { _, err := client.DialContext(dialCtx, "tcp", "target.invalid:443"); dialDone <- err }() + select { + case err := <-dialDone: + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("pending dial error=%v, want context deadline", err) + } + case <-time.After(time.Second): + _ = wire.Close() + t.Fatal("pending dial did not honor its context deadline") + } + closed := make(chan error, 1) + go func() { closed <- client.Close() }() + select { + case err := <-closed: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + _ = wire.Close() + <-closed + t.Fatal("client Close did not interrupt the stalled wire") + } + }) + } +} + +func TestWebH2RetirementKeepsOpeningReservation(t *testing.T) { + wire, h2 := newWebH2StalledWriter(t) + session := &webH2ClientSession{raw: wire, h2: h2, opening: 2} + client := &WebH2Client{sessions: map[*webH2ClientSession]struct{}{session: {}}} + client.releaseSessionReservation(session) + if wire.closes.Load() != 0 || len(client.sessions) != 1 || session.opening != 1 { + t.Fatal("retirement closed a session with another opening stream") + } + client.releaseSessionReservation(session) + if wire.closes.Load() == 0 || len(client.sessions) != 0 || session.opening != 0 { + t.Fatal("last opening reservation did not retire the drained session") + } +} + +func TestWebH2RetirementPreservesSiblingsAndHalfClose(t *testing.T) { + target := startWebTCPEcho(t) + client := newWebH2DeadlineTestClient(t, transport.DialFunc((&net.Dialer{}).DialContext)) + first := dialWebH2DeadlineTestConn(t, client, target) + sibling := dialWebH2DeadlineTestConn(t, client, target) + client.mu.Lock() + old := client.current + client.mu.Unlock() + old.h2.SetDoNotReuse() + fresh := dialWebH2DeadlineTestConn(t, client, target) + assertH2Ownership(t, client, old, 2, 2) + assertWebSessionSiblingEcho(t, first) + assertWebSessionSiblingEcho(t, sibling) + if err := first.Close(); err != nil { + t.Fatal(err) + } + assertH2Ownership(t, client, old, 1, 2) + assertWebSessionSiblingEcho(t, sibling) + if err := sibling.(interface{ CloseWrite() error }).CloseWrite(); err != nil { + t.Fatal(err) + } + assertWebEOFBody(t, sibling, nil) + // Neither upload FIN nor response EOF releases the caller's net.Conn. + assertH2Ownership(t, client, old, 1, 2) + if err := sibling.Close(); err != nil { + t.Fatal(err) + } + assertH2Ownership(t, client, old, 0, 1) + assertWebSessionSiblingEcho(t, fresh) +} + +func TestWebH2DeadlineReleasesRetiredStreamExactlyOnce(t *testing.T) { + target := startWebTCPEcho(t) + client := newWebH2DeadlineTestClient(t, transport.DialFunc((&net.Dialer{}).DialContext)) + oldConn := dialWebH2DeadlineTestConn(t, client, target) + client.mu.Lock() + old := client.current + client.mu.Unlock() + old.h2.SetDoNotReuse() + fresh := dialWebH2DeadlineTestConn(t, client, target) + assertH2Ownership(t, client, old, 1, 2) + if err := oldConn.SetDeadline(time.Now()); err != nil { + t.Fatal(err) + } + var data [1]byte + _, _ = oldConn.Read(data[:]) + // Close joins any concurrent deadline callback, including onClose. + if err := oldConn.Close(); err != nil { + t.Fatal(err) + } + if err := oldConn.Close(); err != nil { + t.Fatal(err) + } + assertH2Ownership(t, client, old, 0, 1) + assertWebSessionSiblingEcho(t, fresh) +} + +func TestWebH2RetiredStreamCloseRacesClientClose(t *testing.T) { + for range 10 { + t.Run("close", func(t *testing.T) { + client := newWebH2DeadlineTestClient(t, transport.DialFunc((&net.Dialer{}).DialContext)) + target := startWebTCPEcho(t) + conn := dialWebH2DeadlineTestConn(t, client, target) + client.mu.Lock() + old := client.current + client.mu.Unlock() + old.h2.SetDoNotReuse() + _ = dialWebH2DeadlineTestConn(t, client, target) + start, done := make(chan struct{}), make(chan struct{}, 3) + for _, operation := range []func(){ + func() { _ = conn.Close() }, + func() { _ = client.Close() }, + func() { _ = conn.SetDeadline(time.Now()) }, + } { + go func() { <-start; operation(); done <- struct{}{} }() + } + close(start) + for range 3 { + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("stream deadline/Close and client Close deadlocked") + } + } + assertH2Ownership(t, client, old, 0, 0) + }) + } +} + +func assertH2Ownership(t *testing.T, client *WebH2Client, session *webH2ClientSession, active, sessions int) { + t.Helper() + client.mu.Lock() + defer client.mu.Unlock() + if session.opening != 0 || session.active != active || len(client.sessions) != sessions { + t.Fatalf("opening/active/sessions=%d/%d/%d, want 0/%d/%d", session.opening, session.active, len(client.sessions), active, sessions) + } +} + +type webH2StalledWriter struct { + net.Conn + signal atomic.Bool + closes atomic.Int32 + started chan struct{} + once sync.Once +} + +func (c *webH2StalledWriter) Write(p []byte) (int, error) { + if c.signal.Load() { + c.once.Do(func() { close(c.started) }) + } + return c.Conn.Write(p) +} + +func (c *webH2StalledWriter) Close() error { + c.closes.Add(1) + return c.Conn.Close() +} + +func newWebH2StalledWriter(t *testing.T) (*webH2StalledWriter, *http2.ClientConn) { + t.Helper() + raw, peer := net.Pipe() + wire := &webH2StalledWriter{Conn: raw, started: make(chan struct{})} + t.Cleanup(func() { _ = raw.Close(); _ = peer.Close() }) + initialRead := make(chan error, 1) + go func() { + prefix := make([]byte, len(http2.ClientPreface)) + if _, err := io.ReadFull(peer, prefix); err != nil { + initialRead <- err + return + } + fr := http2.NewFramer(peer, peer) + for range 2 { // Initial SETTINGS and connection WINDOW_UPDATE. + if _, err := fr.ReadFrame(); err != nil { + initialRead <- err + return + } + } + initialRead <- nil + }() + h2, err := (&http2.Transport{}).NewClientConn(wire) + if err != nil { + t.Fatal(err) + } + // Always release the physical writer before closing the H2 wrapper. + t.Cleanup(func() { _ = raw.Close(); _ = h2.Close() }) + if err := <-initialRead; err != nil { + t.Fatal(err) + } + wire.signal.Store(true) + settingsSent := make(chan error, 1) + go func() { settingsSent <- http2.NewFramer(peer, peer).WriteSettings() }() + select { + case <-wire.started: + case <-time.After(time.Second): + t.Fatal("SETTINGS ACK did not reach the blocked transport writer") + } + if err := <-settingsSent; err != nil { + t.Fatal(err) + } + return wire, h2 +} diff --git a/internal/tunnel/web_stream.go b/internal/tunnel/web_stream.go index 4f86654..c26cb71 100644 --- a/internal/tunnel/web_stream.go +++ b/internal/tunnel/web_stream.go @@ -128,6 +128,7 @@ type webH2Conn struct { writeTimer *time.Timer closeOnce sync.Once closeErr error + onClose func() } func newWebH2Conn(reader io.ReadCloser, writer *io.PipeWriter, cancel context.CancelFunc, local, remote net.Addr) *webH2Conn { @@ -171,7 +172,15 @@ func (c *webH2Conn) Close() error { func (c *webH2Conn) closeStream() error { c.closeOnce.Do(func() { c.cancel() - c.closeErr = errors.Join(c.writer.Close(), c.reader.Close()) + writeErr := c.writer.Close() + // The request is now canceled and its pipe cannot accept more data. + // Release session ownership before Body.Close: that operation may need + // the HTTP/2 write lock to return flow-control credit. If this was the + // last retired stream, closing its wire first unblocks that write lock. + if c.onClose != nil { + c.onClose() + } + c.closeErr = errors.Join(writeErr, c.reader.Close()) }) return c.closeErr } diff --git a/internal/tunnel/web_udp.go b/internal/tunnel/web_udp.go index dd0b6dc..3eb1dc2 100644 --- a/internal/tunnel/web_udp.go +++ b/internal/tunnel/web_udp.go @@ -438,7 +438,10 @@ type webUDPPacketConn struct { closed bool sessions map[string]*webUDPClientSession pending map[string]*webUDPPendingSession - once sync.Once + // Reports only a newly verified CONNECT-UDP response, never local enqueue. + // Protected by mu; invoked after releasing mu to avoid a Close lock cycle. + onAuthenticated func(bool) + once sync.Once } func newWebUDPPacketConn(client *WebH3Client) *webUDPPacketConn { @@ -460,6 +463,21 @@ func newWebUDPPacketConn(client *WebH3Client) *webUDPPacketConn { return packet } +func (p *webUDPPacketConn) setAuthenticationObserver(observer func(bool)) { + p.mu.Lock() + p.onAuthenticated = observer + p.mu.Unlock() +} + +func (p *webUDPPacketConn) authenticatedResponse(connected bool) { + p.mu.Lock() + observer := p.onAuthenticated + p.mu.Unlock() + if observer != nil { + observer(connected) + } +} + func (p *webUDPPacketConn) Send(payload []byte, address string) error { if len(payload) > webConnectUDPMaxPayloadSize { return fmt.Errorf("tunnel: CONNECT-UDP payload exceeds %d bytes", webConnectUDPMaxPayloadSize) @@ -555,8 +573,13 @@ func (p *webUDPPacketConn) openSession(target string) (*webUDPClientSession, err defer cancelOpen() opened, err := p.client.openConnectUDPSession(openContext, target) if err != nil { + var connectErr *WebConnectError + if errors.As(err, &connectErr) { + p.authenticatedResponse(false) + } return nil, err } + p.authenticatedResponse(true) sessionContext, cancelSession := context.WithCancel(p.ctx) session := &webUDPClientSession{ packet: p,