From 82a05108c141920fb6123ad93d94d37fdf8cf78c Mon Sep 17 00:00:00 2001 From: cppla Date: Tue, 22 Sep 2026 20:54:15 +0800 Subject: [PATCH] fix: keep HTTP transfers alive and cancel stalled initialization --- docs/DEPLOYMENT.md | 5 +- docs/WEB_COVER.md | 13 + integration/proxy_activity_test.go | 98 ++++++ internal/proxy/config.go | 4 +- internal/proxy/http.go | 32 +- internal/proxy/http_activity_test.go | 168 +++++++++++ internal/proxy/lifecycle.go | 17 +- internal/proxy/relay.go | 32 +- internal/tunnel/web_h2_client.go | 49 ++- internal/tunnel/web_h2_initialization_test.go | 278 ++++++++++++++++++ internal/tunnel/web_handler.go | 23 ++ .../tunnel/web_server_cancellation_test.go | 265 +++++++++++++++++ 12 files changed, 962 insertions(+), 22 deletions(-) create mode 100644 internal/proxy/http_activity_test.go create mode 100644 internal/tunnel/web_h2_initialization_test.go create mode 100644 internal/tunnel/web_server_cancellation_test.go diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index fbce9ec..c183130 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -248,7 +248,10 @@ 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. +indefinitely. Ordinary HTTP forwarding also counts upload/download body +progress as activity, including small buffered responses. Stalled request or +response bodies and blocked writes remain bounded; header and keepalive +timeouts are unchanged. 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 08f46e0..fd57c81 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -42,6 +42,19 @@ finish its response while continuing to receive an independent upload. Other streams on the same connection remain usable. Applications that need to keep uploading after receiving a destination EOF require the H3 stream transport. +Canceled H2 requests and H3 send-side stream errors close the destination directly, +including when both relay workers are blocked on destination I/O. Server or +physical-connection shutdown also releases H3 destinations after a response +FIN. One H3 edge case remains: after a clean response FIN, resetting only that +stream cannot interrupt an already blocked destination upload write through +the current QUIC API. The stream slot can remain occupied until the destination +unblocks or the physical connection closes; a sibling stream can remain usable. +This is not a guarantee that every canceled stream is immediately reclaimed. + +H2's handshake budget covers both TLS negotiation and the initial HTTP/2 +preface/SETTINGS write. Caller cancellation or client shutdown also interrupts +that initialization, before the connection enters the reusable session pool. + There is no `autocar/2` ALPN or AutoCAR binary stream header on these paths. The web ALPNs are `h2`, `h3`, and `http/1.1`. Native and web transports remain separate modes and are not wire-compatible. diff --git a/integration/proxy_activity_test.go b/integration/proxy_activity_test.go index 152e7fd..10f05ca 100644 --- a/integration/proxy_activity_test.go +++ b/integration/proxy_activity_test.go @@ -8,6 +8,8 @@ import ( "io" "net" "net/http" + "net/http/httptest" + "net/url" "strings" "testing" "time" @@ -17,6 +19,102 @@ import ( "github.com/cppla/autocar/internal/tunnel" ) +func TestHTTPForwardWebH2ActiveBodiesSurviveIdleTimeout(t *testing.T) { + const idle = 250 * time.Millisecond + const chunks = 16 + const interval = 50 * time.Millisecond + dialer := startActivityWebH2Client(t) + for _, direction := range []string{"download", "upload"} { + t.Run(direction, func(t *testing.T) { + chunk := []byte("body-chunk\x00\xff") + payload := bytes.Repeat(chunk, chunks) + origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if direction == "upload" { + got, err := io.ReadAll(r.Body) + if err != nil || !bytes.Equal(got, payload) { + http.Error(w, "incomplete upload", http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusNoContent) + return + } + w.Header().Set("Content-Length", fmt.Sprint(len(payload))) + for range chunks { + if _, err := w.Write(chunk); err != nil { + return + } + w.(http.Flusher).Flush() + select { + case <-r.Context().Done(): + return + case <-time.After(interval): + } + } + })) + defer origin.Close() + address := startActivityHTTPProxy(t, dialer, idle) + client := &http.Client{ + Timeout: operationTimeout, + Transport: &http.Transport{Proxy: http.ProxyURL(&url.URL{Scheme: "http", Host: address})}, + } + defer client.CloseIdleConnections() + method := http.MethodGet + var body io.Reader + if direction == "upload" { + method = http.MethodPost + reader, writer := io.Pipe() + defer reader.Close() + defer writer.Close() + body = reader + done := make(chan struct{}) + t.Cleanup(func() { + select { + case <-done: + case <-time.After(operationTimeout): + t.Error("upload writer did not stop") + } + }) + go func() { + defer close(done) + defer writer.Close() + for range chunks { + if _, err := writer.Write(chunk); err != nil { + return + } + time.Sleep(interval) + } + }() + } + request, err := http.NewRequest(method, origin.URL, body) + if err != nil { + t.Fatal(err) + } + if direction == "upload" { + request.ContentLength = int64(len(payload)) + } + response, err := client.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + got, err := io.ReadAll(response.Body) + if err != nil { + t.Fatalf("active %s over H2: %v (%d bytes)", direction, err, len(got)) + } + wantStatus := http.StatusNoContent + if direction == "download" { + wantStatus = http.StatusOK + if !bytes.Equal(got, payload) { + t.Error("download bytes differ") + } + } + if response.StatusCode != wantStatus { + t.Fatalf("active %s over H2: status %d, want %d: %s", direction, response.StatusCode, wantStatus, got) + } + }) + } +} + // 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. diff --git a/internal/proxy/config.go b/internal/proxy/config.go index d432942..8984669 100644 --- a/internal/proxy/config.go +++ b/internal/proxy/config.go @@ -30,8 +30,8 @@ type Config struct { HandshakeTimeout time.Duration DialTimeout 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. + // tunnel or an HTTP body transfer. Upload/download progress keeps reads + // alive; blocked writes retain an independent timeout. 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 diff --git a/internal/proxy/http.go b/internal/proxy/http.go index a8e360f..efbdd48 100644 --- a/internal/proxy/http.go +++ b/internal/proxy/http.go @@ -30,6 +30,8 @@ type HTTPServer struct { transport *http.Transport } +type httpConnContextKey struct{} + // NewHTTPServer validates cfg and creates an HTTP forward proxy. func NewHTTPServer(cfg Config) (*HTTPServer, error) { normalized, err := normalizeConfig(cfg) @@ -56,6 +58,12 @@ func NewHTTPServer(cfg Config) (*HTTPServer, error) { ReadHeaderTimeout: normalized.handshakeTimeout, IdleTimeout: normalized.idleTimeout, MaxHeaderBytes: 64 << 10, + ConnContext: func(ctx context.Context, conn net.Conn) context.Context { + if tracked, ok := conn.(*trackedConn); ok { + return context.WithValue(ctx, httpConnContextKey{}, tracked) + } + return ctx + }, ConnState: func(conn net.Conn, state http.ConnState) { tracked, ok := conn.(*trackedConn) if !ok { @@ -272,7 +280,14 @@ func (s *HTTPServer) serveForward(w http.ResponseWriter, r *http.Request) { } } w.WriteHeader(response.StatusCode) - _, _ = io.Copy(w, response.Body) + var body io.Reader = response.Body + if conn, ok := r.Context().Value(httpConnContextKey{}).(*trackedConn); ok { + // net/http waits for disconnects in a background read after the + // request body ends. Origin data is activity even while the response + // writer buffers small chunks and the client sends nothing more. + body = &httpResponseActivityReader{Reader: body, conn: conn} + } + _, _ = io.Copy(w, body) for key, values := range response.Trailer { if isHopByHopHeader(key) { continue @@ -281,6 +296,21 @@ func (s *HTTPServer) serveForward(w http.ResponseWriter, r *http.Request) { } } +type httpResponseActivityReader struct { + io.Reader + conn *trackedConn +} + +func (r *httpResponseActivityReader) Read(p []byte) (int, error) { + n, err := r.Reader.Read(p) + if n > 0 { + if activityErr := r.conn.refreshReadActivity(); err == nil { + err = activityErr + } + } + return n, err +} + func validateAbsoluteTarget(target *url.URL) error { host := target.Hostname() if host == "" { diff --git a/internal/proxy/http_activity_test.go b/internal/proxy/http_activity_test.go new file mode 100644 index 0000000..9a79695 --- /dev/null +++ b/internal/proxy/http_activity_test.go @@ -0,0 +1,168 @@ +package proxy + +import ( + "bytes" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "testing" + "time" +) + +func TestHTTPForwardActiveBodiesSurviveIdleTimeout(t *testing.T) { + const idle = 200 * time.Millisecond + const chunks = 16 + const interval = idle / 5 + for _, direction := range []string{"download", "upload"} { + t.Run(direction, func(t *testing.T) { + chunk := bytes.Repeat([]byte("a"), 16) + origin := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if direction == "upload" { + got, err := io.ReadAll(r.Body) + if err != nil || !bytes.Equal(got, bytes.Repeat(chunk, chunks)) { + http.Error(w, "incomplete upload", http.StatusBadRequest) + return + } + w.WriteHeader(http.StatusNoContent) + return + } + w.Header().Set("Content-Length", fmt.Sprint(chunks*len(chunk))) + for range chunks { + if _, err := w.Write(chunk); err != nil { + return + } + w.(http.Flusher).Flush() + select { + case <-r.Context().Done(): + return + case <-time.After(interval): + } + } + })) + defer origin.Close() + server, proxyURL, stop := startHTTPProxy(t, Config{Dialer: directDialer(), IdleTimeout: idle}) + defer stop(server) + client := proxyHTTPClient(t, proxyURL, nil) + defer client.CloseIdleConnections() + method := http.MethodGet + var body io.Reader + if direction == "upload" { + method = http.MethodPost + reader, writer := io.Pipe() + defer reader.Close() + defer writer.Close() + body = reader + done := make(chan struct{}) + t.Cleanup(func() { + select { + case <-done: + case <-time.After(time.Second): + t.Error("upload writer did not stop") + } + }) + go func() { + defer close(done) + defer writer.Close() + for range chunks { + if _, err := writer.Write(chunk); err != nil { + return + } + time.Sleep(interval) + } + }() + } + request, err := http.NewRequest(method, origin.URL, body) + if err != nil { + t.Fatal(err) + } + if direction == "upload" { + request.ContentLength = int64(chunks * len(chunk)) + } + response, err := client.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + got, err := io.ReadAll(response.Body) + if err != nil { + t.Fatalf("read active %s: %v (%d bytes)", direction, err, len(got)) + } + wantStatus := http.StatusNoContent + if direction == "download" { + wantStatus = http.StatusOK + if !bytes.Equal(got, bytes.Repeat(chunk, chunks)) { + t.Errorf("active download = %d bytes, want %d", len(got), chunks*len(chunk)) + } + } + if response.StatusCode != wantStatus { + t.Fatalf("active %s status = %d, want %d: %s", direction, response.StatusCode, wantStatus, got) + } + }) + } +} + +func TestHTTPActivityPreservesExternalReadDeadline(t *testing.T) { + local, peer := net.Pipe() + defer local.Close() + defer peer.Close() + conn := &trackedConn{Conn: local} + conn.setActivityTimeout(time.Hour) + // net/http uses this past deadline to stop its disconnect reader before + // hijacking a CONNECT. Neither forwarded body data nor writes may undo it. + if err := conn.SetReadDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatal(err) + } + if err := conn.refreshReadActivity(); err != nil { + t.Fatal(err) + } + done := make(chan error, 1) + go func() { + _, err := conn.Read(make([]byte, 1)) + done <- err + }() + select { + case err := <-done: + var timeout net.Error + if !errors.As(err, &timeout) || !timeout.Timeout() { + t.Fatalf("read error = %v, want preserved deadline", err) + } + case <-time.After(time.Second): + t.Fatal("activity extended net/http's external read deadline") + } +} + +func TestHTTPActivityDoesNotExtendBlockedWrite(t *testing.T) { + local, peer := net.Pipe() + defer local.Close() + defer peer.Close() + conn := &trackedConn{Conn: local} + conn.setActivityTimeout(100 * time.Millisecond) + done := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("blocked")) + done <- err + }() + ticker := time.NewTicker(10 * time.Millisecond) + defer ticker.Stop() + limit := time.NewTimer(time.Second) + defer limit.Stop() + for { + select { + case <-ticker.C: + if err := conn.refreshReadActivity(); err != nil { + t.Fatal(err) + } + case err := <-done: + var timeout net.Error + if !errors.As(err, &timeout) || !timeout.Timeout() { + t.Fatalf("write error = %v, want stall timeout", err) + } + return + case <-limit.C: + t.Fatal("read activity kept a blocked client write alive") + } + } +} diff --git a/internal/proxy/lifecycle.go b/internal/proxy/lifecycle.go index 514d737..aa93496 100644 --- a/internal/proxy/lifecycle.go +++ b/internal/proxy/lifecycle.go @@ -111,7 +111,22 @@ func (c *trackedConn) Write(p []byte) (int, error) { return 0, err } } - return c.Conn.Write(p) + n, err := c.Conn.Write(p) + if n > 0 { + if activityErr := c.refreshReadActivity(); err == nil { + err = activityErr + } + } + return n, err +} + +func (c *trackedConn) refreshReadActivity() error { + c.deadlineMu.Lock() + defer c.deadlineMu.Unlock() + if timeout := time.Duration(c.activityTimeout.Load()); timeout > 0 { + return c.Conn.SetReadDeadline(earlierDeadline(time.Now().Add(timeout), c.externalRead)) + } + return nil } // SetDeadline records deadlines imposed by net/http. Activity deadlines are diff --git a/internal/proxy/relay.go b/internal/proxy/relay.go index 8f09746..7bfaf3c 100644 --- a/internal/proxy/relay.go +++ b/internal/proxy/relay.go @@ -13,25 +13,39 @@ var relayBufferPool = sync.Pool{New: func() any { return &buffer }} -// activityConn applies an inactivity timeout to every blocking read and -// write. It is used for HTTP origin connections, whose request/response body -// phase is otherwise not covered by net/http's header and keepalive timeouts. +// activityConn bounds HTTP origin reads by connection inactivity and writes +// by their own stall timeout. Upload progress keeps the concurrent response +// read alive even when the origin waits for the full request body to respond. type activityConn struct { net.Conn - timeout time.Duration + timeout time.Duration + readDeadlineMu sync.Mutex } func (c *activityConn) Read(p []byte) (int, error) { - if c.timeout > 0 { - if err := c.Conn.SetReadDeadline(time.Now().Add(c.timeout)); err != nil { - return 0, err - } + if err := c.refreshReadActivity(); err != nil { + return 0, err } return c.Conn.Read(p) } func (c *activityConn) Write(p []byte) (int, error) { - return writeWithStallDeadline(c.Conn, p, c.timeout) + n, err := writeWithStallDeadline(c.Conn, p, c.timeout) + if n > 0 { + if activityErr := c.refreshReadActivity(); err == nil { + err = activityErr + } + } + return n, err +} + +func (c *activityConn) refreshReadActivity() error { + if c.timeout <= 0 { + return nil + } + c.readDeadlineMu.Lock() + defer c.readDeadlineMu.Unlock() + return c.Conn.SetReadDeadline(time.Now().Add(c.timeout)) } // Bound only the pending write. In particular, H2 streams implement deadlines diff --git a/internal/tunnel/web_h2_client.go b/internal/tunnel/web_h2_client.go index 5bc7431..1b99469 100644 --- a/internal/tunnel/web_h2_client.go +++ b/internal/tunnel/web_h2_client.go @@ -39,7 +39,7 @@ type WebH2Client struct { tlsConfig *tls.Config fingerprint FingerprintProfile handshakeTimeout time.Duration - dialer net.Dialer + dialer transport.Dialer auth *webAuthSigner claims webAuthClaims transport *http2.Transport @@ -141,7 +141,7 @@ func newWebH2ClientWithSigner(config WebH2ClientConfig, auth *webAuthSigner, cla tlsConfig: tlsConfig, fingerprint: fingerprint, handshakeTimeout: handshakeTimeout, - dialer: net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}, + dialer: &net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}, auth: auth, claims: claims, transport: &http2.Transport{ @@ -477,32 +477,65 @@ func (c *WebH2Client) openSession(ctx context.Context) (*webH2ClientSession, err 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. + initializationCtx, initializationCancel := context.WithTimeout(dialCtx, c.handshakeTimeout) + defer initializationCancel() + rawClosed := make(chan struct{}) + stopRawClose := context.AfterFunc(initializationCtx, func() { + _ = raw.Close() + close(rawClosed) + }) + watcherDetached := false + defer func() { + if !watcherDetached && !stopRawClose() { + <-rawClosed + } + }() tlsConn, err := newWebH2TLSClientConn(raw, c.tlsConfig, c.fingerprint, c.utlsSessionCache) if err != nil { _ = raw.Close() return nil, err } - handshakeCtx, handshakeCancel := context.WithTimeout(dialCtx, c.handshakeTimeout) - err = tlsConn.HandshakeContext(handshakeCtx) - handshakeCancel() + err = tlsConn.HandshakeContext(initializationCtx) if err != nil { _ = raw.Close() + if cause := context.Cause(initializationCtx); cause != nil { + err = cause + } return nil, fmt.Errorf("tunnel: web-cover HTTP/2 TLS handshake: %w", err) } state := tlsConn.ConnectionState() if state.Version != tls.VersionTLS13 { - _ = tlsConn.Close() + _ = raw.Close() return nil, errors.New("tunnel: web-cover HTTP/2 connection did not negotiate TLS 1.3") } if state.NegotiatedProtocol != webH2ALPN { - _ = tlsConn.Close() + _ = raw.Close() return nil, errors.New("tunnel: web-cover HTTP/2 connection did not negotiate h2") } clientConn, err := c.transport.NewClientConn(tlsConn) if err != nil { - _ = tlsConn.Close() + _ = raw.Close() + if cause := context.Cause(initializationCtx); cause != nil { + err = cause + } return nil, fmt.Errorf("tunnel: initialize web-cover HTTP/2 connection: %w", err) } + // Stop the watcher before the deferred initialization cancellation. If it + // already started, join it and reject the connection instead of handing a + // caller a session whose wire is concurrently being closed. + watcherDetached = stopRawClose() + if !watcherDetached { + <-rawClosed + } + if cause := context.Cause(initializationCtx); cause != nil { + _ = raw.Close() + _ = clientConn.Close() + return nil, fmt.Errorf("tunnel: initialize web-cover HTTP/2 connection: %w", cause) + } return &webH2ClientSession{ raw: raw, conn: tlsConn, diff --git a/internal/tunnel/web_h2_initialization_test.go b/internal/tunnel/web_h2_initialization_test.go new file mode 100644 index 0000000..92e0d1d --- /dev/null +++ b/internal/tunnel/web_h2_initialization_test.go @@ -0,0 +1,278 @@ +package tunnel + +import ( + "context" + "crypto/tls" + "crypto/x509" + "encoding/binary" + "errors" + "net" + "net/http" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +func TestWebH2PostTLSInitializationHonorsCancellation(t *testing.T) { + for _, profile := range []FingerprintProfile{FingerprintNative, FingerprintChrome133} { + t.Run(string(profile), func(t *testing.T) { + for _, mode := range []string{"caller_cancel", "caller_deadline", "client_close", "initialization_deadline"} { + t.Run(mode, func(t *testing.T) { + serverTLS, clientTLS := testTLSConfigs(t) + serverTLS.NextProtos = []string{webH2ALPN} + raw, peer := net.Pipe() + wire := &webH2InitializationWire{Conn: raw, initializing: make(chan struct{}), closed: make(chan struct{})} + clientTLS.VerifyPeerCertificate = func(_ [][]byte, _ [][]*x509.Certificate) error { + wire.verified.Store(true) + return nil + } + handshakeTimeout := 2 * time.Second + if mode == "initialization_deadline" { + handshakeTimeout = 250 * time.Millisecond + } + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: "127.0.0.1:443", Token: testToken, TLSConfig: clientTLS, + FingerprintProfile: profile, HandshakeTimeout: handshakeTimeout, + }) + if err != nil { + 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() }) + serverDone := make(chan error, 1) + go func() { serverDone <- tls.Server(peer, serverTLS).HandshakeContext(t.Context()) }() + ctx, cancel := context.WithCancelCause(context.Background()) + defer cancel(context.Canceled) + var dialCtx context.Context = ctx + if mode == "caller_deadline" { + var cancelDeadline context.CancelFunc + dialCtx, cancelDeadline = context.WithTimeout(ctx, 250*time.Millisecond) + defer cancelDeadline() + } + dialDone := make(chan error, 1) + go func() { + conn, err := client.DialContext(dialCtx, "tcp", "target.invalid:443") + if conn != nil { + _ = conn.Close() + } + dialDone <- err + }() + select { + case <-wire.initializing: + case err := <-dialDone: + t.Fatalf("dial finished before the post-TLS write: %v", err) + 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) + } + client.mu.Lock() + registered := len(client.sessions) + client.mu.Unlock() + if registered != 0 { + t.Fatalf("initializing sessions already registered: %d", registered) + } + var want error + switch mode { + case "caller_cancel": + want = errors.New("caller stopped H2 initialization") + cancel(want) + case "caller_deadline": + want = context.DeadlineExceeded + <-dialCtx.Done() + case "client_close": + want = context.Canceled + if err := client.Close(); err != nil { + t.Fatal(err) + } + case "initialization_deadline": + want = context.DeadlineExceeded + } + select { + case err := <-dialDone: + if !errors.Is(err, want) { + 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") + } + select { + case <-wire.closed: + default: + t.Fatal("canceled initializer left its raw connection open") + } + client.mu.Lock() + registered = len(client.sessions) + client.mu.Unlock() + if registered != 0 || len(client.dialGate) != 0 { + t.Fatalf("canceled initializer left sessions/gate=%d/%d", registered, len(client.dialGate)) + } + }) + } + }) + } +} + +func TestWebH2InitializationWatcherDetachesAfterSuccess(t *testing.T) { + for _, profile := range []FingerprintProfile{FingerprintNative, FingerprintChrome133} { + t.Run(string(profile), func(t *testing.T) { + serverTLS, clientTLS := testTLSConfigs(t) + server := startWebH2TestServer(t, WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + Cover: http.NotFoundHandler(), Dialer: &net.Dialer{}, + }) + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: server.Addr().String(), Token: testToken, TLSConfig: clientTLS, + FingerprintProfile: profile, HandshakeTimeout: 250 * time.Millisecond, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + dialCtx, cancel := context.WithCancel(context.Background()) + defer cancel() + conn, err := client.DialContext(dialCtx, "tcp", startWebTCPEcho(t)) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + cancel() + // Both caller cancellation and the elapsed initialization timeout + // must leave ownership with the returned stream and warm session. + time.Sleep(300 * time.Millisecond) + assertWebSessionSiblingEcho(t, conn) + }) + } +} + +func TestWebH2CloseAtInitializationHandoff(t *testing.T) { + for _, profile := range []FingerprintProfile{FingerprintNative, FingerprintChrome133} { + t.Run(string(profile), func(t *testing.T) { + for range 10 { + t.Run("close", func(t *testing.T) { + serverTLS, clientTLS := testTLSConfigs(t) + serverTLS.NextProtos = []string{webH2ALPN} + raw, peer := net.Pipe() + wire := &webH2InitializationWire{Conn: raw, initializing: make(chan struct{}), closed: make(chan struct{})} + clientTLS.VerifyPeerCertificate = func(_ [][]byte, _ [][]*x509.Certificate) error { + wire.verified.Store(true) + return nil + } + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: "127.0.0.1:443", Token: testToken, TLSConfig: clientTLS, + FingerprintProfile: profile, HandshakeTimeout: time.Second, + }) + if err != nil { + 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{}) + go func() { + defer close(serverDone) + server := tls.Server(peer, serverTLS) + if server.HandshakeContext(t.Context()) == nil { + var data [1]byte + _, _ = server.Read(data[:]) + } + }() + dialDone := make(chan error, 1) + go func() { + conn, err := client.DialContext(context.Background(), "tcp", "target.invalid:443") + if conn != nil { + _ = conn.Close() + } + dialDone <- err + }() + select { + case err := <-dialDone: + if err == nil { + t.Fatal("closed client published a cold initialized session") + } + case <-time.After(2 * time.Second): + _ = raw.Close() + <-dialDone + t.Fatal("close during initialization handoff did not finish") + } + <-serverDone + select { + case <-wire.closed: + default: + t.Fatal("closed initializer left its raw connection open") + } + client.mu.Lock() + registered, current, closed := len(client.sessions), client.current, client.closed + client.mu.Unlock() + if !closed { + t.Fatal("test did not reach Close after the initial H2 write") + } + if registered != 0 || current != nil { + t.Fatal("closed client retained a session after initialization handoff") + } + }) + } + }) + } +} + +// 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. +// Signaling that write lets tests cancel precisely after TLS, without sleeps +// or runtime-stack inspection to locate the initialization boundary. +type webH2InitializationWire struct { + net.Conn + verified atomic.Bool + finished atomic.Bool + initializing chan struct{} + closed chan struct{} + startOnce sync.Once + closeOnce sync.Once + handoffOnce sync.Once + onInitialized func() +} + +func (c *webH2InitializationWire) Write(p []byte) (int, error) { + initialWrite := c.finished.Load() + if initialWrite { + c.startOnce.Do(func() { close(c.initializing) }) + } + n, err := c.Conn.Write(p) + if err == nil && c.verified.Load() && containsH2TestTLSApplicationRecord(p[:n]) { + c.finished.Store(true) + } + if initialWrite && err == nil && c.onInitialized != nil { + c.handoffOnce.Do(c.onInitialized) + } + return n, err +} + +func (c *webH2InitializationWire) Close() error { + err := c.Conn.Close() + c.closeOnce.Do(func() { close(c.closed) }) + return err +} + +func containsH2TestTLSApplicationRecord(data []byte) bool { + for len(data) >= 5 { + size := 5 + int(binary.BigEndian.Uint16(data[3:5])) + if size > len(data) { + return false + } + if data[0] == 23 { // TLS record type application_data (encrypted TLS 1.3). + return true + } + data = data[size:] + } + return false +} diff --git a/internal/tunnel/web_handler.go b/internal/tunnel/web_handler.go index fb5e99a..a68dc04 100644 --- a/internal/tunnel/web_handler.go +++ b/internal/tunnel/web_handler.go @@ -10,6 +10,7 @@ import ( "sync" "time" + "github.com/apernet/quic-go" "github.com/apernet/quic-go/http3" ) @@ -126,6 +127,23 @@ func (h *webTunnelHandler) serveH3Connect( remote, _ := r.Context().Value(http3.RemoteAddrContextKey).(net.Addr) conn := newWebH3Conn(stream, local, remote) _ = conn.SetDeadline(time.Time{}) + // Both copies may be blocked in the destination (one Read, one Write), + // so closing only the QUIC stream cannot wake either of them. A normal + // response FIN also cancels the stream context, however: preserve that + // half-close and abort the destination only for an actual stream error. + stopStream := context.AfterFunc(stream.Context(), func() { + if !errors.Is(context.Cause(stream.Context()), context.Canceled) { + _ = upstream.Close() + } + }) + defer stopStream() + // After a clean response FIN the stream context is already canceled and + // cannot report a later connection failure. Keep that shutdown signal + // separate so a still-running upload is also released by server Close. + if physical, ok := r.Context().Value(webH3ConnectionContextKey{}).(*quic.Conn); ok { + stopConnection := context.AfterFunc(physical.Context(), func() { _ = upstream.Close() }) + defer stopConnection() + } relay(conn, upstream) } @@ -135,6 +153,11 @@ func (h *webTunnelHandler) serveH2Connect( upstream net.Conn, authentication *webRequestAuthentication, ) { + // The request context is canceled on stream reset or server shutdown, + // unlike an orderly upload FIN. Closing the destination directly also + // releases a copy already blocked in its Write, not just a body Read. + stop := context.AfterFunc(r.Context(), func() { _ = upstream.Close() }) + defer stop() flusher, ok := w.(http.Flusher) if !ok { h.writeAuthenticatedError(w, authentication, http.StatusBadGateway) diff --git a/internal/tunnel/web_server_cancellation_test.go b/internal/tunnel/web_server_cancellation_test.go new file mode 100644 index 0000000..1da3360 --- /dev/null +++ b/internal/tunnel/web_server_cancellation_test.go @@ -0,0 +1,265 @@ +package tunnel + +import ( + "bytes" + "context" + "fmt" + "io" + "net" + "net/http" + "sync" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +func TestWebServerCancellationReleasesStalledDestination(t *testing.T) { + for _, mode := range []string{"h2", "h3"} { + actions := []string{"close_stream", "close_server"} + if mode == "h3" { + actions = append(actions, "close_server_after_response_fin") + } + for _, action := range actions { + t.Run(mode+"/"+action, func(t *testing.T) { + upstream, target := net.Pipe() + defer target.Close() + tracked := &webServerStalledTarget{ + Conn: upstream, writing: make(chan struct{}), closed: make(chan struct{}), + } + defer tracked.Close() + admission, err := NewStreamAdmission(2) + if err != nil { + t.Fatal(err) + } + outbound := transport.DialFunc(func(ctx context.Context, network, address string) (net.Conn, error) { + if address == "stall.example:443" { + if action == "close_server_after_response_fin" { + return &webServerResponseEOF{tracked}, nil + } + return tracked, nil + } + return (&net.Dialer{}).DialContext(ctx, network, address) + }) + client, closeServer := newWebServerCancellationClient(t, mode, outbound, admission) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + sibling, err := client.DialContext(ctx, "tcp", startWebTCPEcho(t)) + if err != nil { + t.Fatal(err) + } + defer sibling.Close() + conn, err := client.DialContext(ctx, "tcp", "stall.example:443") + if err != nil { + t.Fatal(err) + } + defer conn.Close() + if err := conn.SetDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatal(err) + } + if action == "close_server_after_response_fin" { + if reply, err := io.ReadAll(conn); err != nil || len(reply) != 0 { + t.Fatalf("initial response FIN: %q, %v", reply, err) + } + } + if _, err := io.WriteString(conn, "x"); err != nil { + t.Fatal(err) + } + select { + case <-tracked.writing: + case <-time.After(time.Second): + t.Fatal("upload never reached the blocking destination Write") + } + // Both relay copies now wait on the destination, not on the + // HTTP stream. A reset must actively close that destination. + if action == "close_stream" { + _ = conn.Close() + } else { + closed := make(chan error, 1) + go func() { closed <- closeServer() }() + select { + case <-closed: + case <-time.After(time.Second): + _ = tracked.Close() + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("server Close stayed blocked after forced target cleanup") + } + t.Fatal("server Close waited for a stalled destination") + } + } + select { + case <-tracked.closed: + case <-time.After(time.Second): + t.Fatal("cancellation left destination Read and Write blocked") + } + if action == "close_stream" { + awaitWebServerAdmission(t, admission, 1) + assertWebSessionSiblingEcho(t, sibling) + _ = sibling.Close() + } + awaitWebServerAdmission(t, admission, 0) + }) + } + } +} + +// QUIC's send-stream context is canceled by an ordinary response FIN too. +// That signal must not abort the client's still-live upload direction. +func TestWebH3ServerResponseFINPreservesUpload(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + wantUpload := []byte("upload after response EOF") + wantResponse := []byte("response before upload") + result := make(chan error, 1) + done := make(chan struct{}) + go func() { + defer close(done) + 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(3 * time.Second)); err != nil { + return err + } + if _, err := conn.Write(wantResponse); err != nil { + return err + } + if err := conn.(*net.TCPConn).CloseWrite(); err != nil { + return err + } + got, err := io.ReadAll(conn) + if err != nil { + return err + } + if !bytes.Equal(got, wantUpload) { + return fmt.Errorf("upload = %q, want %q", got, wantUpload) + } + return nil + }() + }() + t.Cleanup(func() { + cancel() + _ = listener.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Error("half-closed target did not stop") + } + }) + admission, err := NewStreamAdmission(1) + if err != nil { + t.Fatal(err) + } + client, _ := newWebServerCancellationClient(t, "h3", &net.Dialer{}, admission) + dialCtx, cancelDial := context.WithTimeout(ctx, 2*time.Second) + defer cancelDial() + conn, err := client.DialContext(dialCtx, "tcp", listener.Addr().String()) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + if err := conn.SetDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatal(err) + } + got, err := io.ReadAll(conn) + if err != nil || !bytes.Equal(got, wantResponse) { + t.Fatalf("response before upload = %q, err = %v", got, err) + } + if _, err := conn.Write(wantUpload); err != nil { + t.Fatalf("upload after response FIN: %v", err) + } + if err := conn.(interface{ CloseWrite() error }).CloseWrite(); err != nil { + t.Fatal(err) + } + select { + case err := <-result: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("upload did not finish after server response FIN") + } + awaitWebServerAdmission(t, admission, 0) +} + +func newWebServerCancellationClient(t *testing.T, mode string, outbound transport.Dialer, admission *StreamAdmission) (transport.Dialer, func() error) { + t.Helper() + serverTLS, clientTLS := testTLSConfigs(t) + if mode == "h2" { + server := startWebH2TestServer(t, WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: webTestToken, TLSConfig: serverTLS, + Cover: http.NotFoundHandler(), Dialer: outbound, StreamAdmission: admission, + }) + client, err := NewWebH2Client(WebH2ClientConfig{ + ServerAddress: server.Addr().String(), Token: webTestToken, TLSConfig: clientTLS, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + return client, server.Close + } + server, err := ListenWebH3(WebH3ServerConfig{ + Address: "127.0.0.1:0", Token: webTestToken, TLSConfig: serverTLS, + Cover: http.NotFoundHandler(), Dialer: outbound, StreamAdmission: admission, + }) + if err != nil { + t.Fatal(err) + } + serveWebH3ForTest(t, server) + client, err := NewWebH3Client(WebH3ClientConfig{ + ServerAddress: server.Addr().String(), Token: webTestToken, TLSConfig: clientTLS, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + return client, server.Close +} + +func awaitWebServerAdmission(t *testing.T, admission *StreamAdmission, want int) { + t.Helper() + timeout := time.NewTimer(time.Second) + defer timeout.Stop() + ticker := time.NewTicker(5 * time.Millisecond) + defer ticker.Stop() + for len(admission.sem) != want { + select { + case <-ticker.C: + case <-timeout.C: + t.Fatalf("occupied stream slots = %d, want %d", len(admission.sem), want) + } + } +} + +type webServerStalledTarget struct { + net.Conn + writing, closed chan struct{} + writeOnce, closeOnce sync.Once +} + +func (c *webServerStalledTarget) Write(p []byte) (int, error) { + c.writeOnce.Do(func() { close(c.writing) }) + return c.Conn.Write(p) +} + +func (c *webServerStalledTarget) Close() error { + var err error + c.closeOnce.Do(func() { err = c.Conn.Close(); close(c.closed) }) + return err +} + +// The target has finished sending but still accepts uploads, which can stall. +type webServerResponseEOF struct{ *webServerStalledTarget } + +func (*webServerResponseEOF) Read([]byte) (int, error) { return 0, io.EOF }