diff --git a/cmd/autocar/preflight.go b/cmd/autocar/preflight.go index 104ac37..7e3ac7d 100644 --- a/cmd/autocar/preflight.go +++ b/cmd/autocar/preflight.go @@ -86,13 +86,16 @@ func checkServerAddresses(mode, udp, tcp string, disableFallback bool) error { // These bounds mirror the transport constructors, which cannot be invoked // for a server without binding sockets. Use them on both the check and normal // startup paths, before any listener is prepared. -func validateServerLocalLimits(token string, handshakeTimeout, dialTimeout time.Duration, maxUpload, maxDownload uint64) error { +func validateServerLocalLimits(token string, handshakeTimeout, dialTimeout, destinationWriteTimeout time.Duration, maxUpload, maxDownload uint64) error { if len(token) < protocol.MinTokenLength || len(token) > protocol.MaxTokenLength { return fmt.Errorf("tunnel: token length must be between %d and %d bytes", protocol.MinTokenLength, protocol.MaxTokenLength) } if handshakeTimeout < 0 || dialTimeout < 0 { return errors.New("--handshake-timeout and --dial-timeout must not be negative") } + if destinationWriteTimeout < 0 { + return errors.New("--destination-write-timeout must not be negative; zero uses the 5m default") + } if maxUpload > protocol.MaxRate || maxDownload > protocol.MaxRate { return errors.New("tunnel: pacing rate exceeds protocol maximum") } diff --git a/cmd/autocar/preflight_test.go b/cmd/autocar/preflight_test.go index 8176bd9..5bd4cd9 100644 --- a/cmd/autocar/preflight_test.go +++ b/cmd/autocar/preflight_test.go @@ -237,6 +237,7 @@ func TestPreflightServerRejectsInvalidConfiguration(t *testing.T) { {"CIDR policy", "invalid denied CIDR", []string{"--deny-cidrs", "not-a-CIDR"}}, {"dial timeout", "--dial-timeout", []string{"--dial-timeout", "-1s"}}, {"handshake timeout", "--handshake-timeout", []string{"--handshake-timeout", "-1s"}}, + {"destination write timeout", "--destination-write-timeout", []string{"--destination-write-timeout", "-1s"}}, {"connections", "--max-connections", []string{"--max-connections", "0"}}, {"UDP sessions", "--max-client-udp-sessions", []string{"--max-udp-sessions", "1"}}, {"rate bound", "protocol maximum", []string{"--pacing", "fixed-rate", "--max-upload-mbps", "8000001", "--max-download-mbps", "1"}}, diff --git a/cmd/autocar/server.go b/cmd/autocar/server.go index ab96c8e..f0e199d 100644 --- a/cmd/autocar/server.go +++ b/cmd/autocar/server.go @@ -56,6 +56,7 @@ func runServer(parent context.Context, args []string) error { allowClientRates := fs.Bool("allow-client-rates", false, "native protocol: allow authenticated clients to request rates within server maxima") dialTimeout := fs.Duration("dial-timeout", 4*time.Second, "remote destination dial timeout") handshakeTimeout := fs.Duration("handshake-timeout", 10*time.Second, "authentication and initial stream-open timeout") + destinationWriteTimeout := fs.Duration("destination-write-timeout", 5*time.Minute, "completion timeout for each TCP destination write of at most 32 KiB; zero uses 5m (not an idle timeout)") if err := parseFlagsWithConfig(fs, args); err != nil { return err } @@ -128,7 +129,7 @@ func runServer(parent context.Context, args []string) error { if err != nil { return err } - if err := validateServerLocalLimits(token, *handshakeTimeout, *dialTimeout, maxUpload, maxDownload); err != nil { + if err := validateServerLocalLimits(token, *handshakeTimeout, *dialTimeout, *destinationWriteTimeout, maxUpload, maxDownload); err != nil { return err } certificate, err := security.LoadKeyPair(*certFile, *keyFile) @@ -188,22 +189,23 @@ func runServer(parent context.Context, args []string) error { } if serverProtocol == "web" { webServer, err := tunnel.ListenWeb(tunnel.WebServerConfig{ - TCPAddress: *tcpListen, - UDPAddress: *listen, - Token: token, - TLSConfig: tlsConfig, - Cover: coverHandler, - Dialer: safeDialer, - HandshakeTimeout: *handshakeTimeout, - DialTimeout: *dialTimeout, - MaxConcurrentStreams: *maxStreams, - StreamAdmission: streamAdmission, - MaxConnections: *maxConnections, - MaxClientConnections: *maxClientConnections, - MaxUDPSessions: *maxUDPSessions, - MaxClientUDPSessions: *maxClientUDPSessions, - MaxUDPDestinations: *maxUDPDestinations, - UDPReceiveQueue: *udpReceiveQueue, + TCPAddress: *tcpListen, + UDPAddress: *listen, + Token: token, + TLSConfig: tlsConfig, + Cover: coverHandler, + Dialer: safeDialer, + HandshakeTimeout: *handshakeTimeout, + DialTimeout: *dialTimeout, + DestinationWriteTimeout: *destinationWriteTimeout, + MaxConcurrentStreams: *maxStreams, + StreamAdmission: streamAdmission, + MaxConnections: *maxConnections, + MaxClientConnections: *maxClientConnections, + MaxUDPSessions: *maxUDPSessions, + MaxClientUDPSessions: *maxClientUDPSessions, + MaxUDPDestinations: *maxUDPDestinations, + UDPReceiveQueue: *udpReceiveQueue, }) if err != nil { return err @@ -229,6 +231,7 @@ func runServer(parent context.Context, args []string) error { AllowClientRates: *allowClientRates, HandshakeTimeout: *handshakeTimeout, DialTimeout: *dialTimeout, + DestinationWriteTimeout: *destinationWriteTimeout, MaxConcurrentStreams: *maxStreams, StreamAdmission: streamAdmission, MaxConnections: *maxConnections, @@ -249,15 +252,16 @@ func runServer(parent context.Context, args []string) error { var tlsServer *tunnel.TLSServer if !*disableFallback { tlsServer, err = tunnel.ListenTLS(tunnel.TLSServerConfig{ - Address: *tcpListen, - Token: token, - TLSConfig: tlsConfig, - Dialer: safeDialer, - HandshakeTimeout: *handshakeTimeout, - DialTimeout: *dialTimeout, - MaxConcurrentStreams: *maxStreams, - StreamAdmission: streamAdmission, - MaxClientConnections: *maxClientFallbackConnections, + Address: *tcpListen, + Token: token, + TLSConfig: tlsConfig, + Dialer: safeDialer, + HandshakeTimeout: *handshakeTimeout, + DialTimeout: *dialTimeout, + DestinationWriteTimeout: *destinationWriteTimeout, + MaxConcurrentStreams: *maxStreams, + StreamAdmission: streamAdmission, + MaxClientConnections: *maxClientFallbackConnections, }) if err != nil { return err diff --git a/cmd/autocar/server_destination_timeout_test.go b/cmd/autocar/server_destination_timeout_test.go new file mode 100644 index 0000000..9fddf2a --- /dev/null +++ b/cmd/autocar/server_destination_timeout_test.go @@ -0,0 +1,104 @@ +package main + +import ( + "context" + "flag" + "io" + "strings" + "testing" + "time" +) + +func TestServerDestinationWriteTimeoutCLIAndCheck(t *testing.T) { + clearPreflightEnvironment(t) + files := newPreflightFiles(t) + dnsCalls := denyPreflightDNS(t) + for _, mode := range []string{"native", "web"} { + for _, value := range []string{"", "0s", "37ms", "8m"} { + t.Run(mode+"/"+value, func(t *testing.T) { + args := append(files.serverArgs(), "--protocol", mode) + if mode == "web" { + args = append(args, "--cover-root", files.cover) + } + if value != "" { + args = append(args, "--destination-write-timeout", value) + } + if err := runServer(context.Background(), args); err != nil { + t.Fatal(err) + } + }) + } + for _, check := range []string{"true", "false"} { + args := append(files.serverArgs(), "--protocol", mode, "--check="+check, "--destination-write-timeout=-1s") + if mode == "web" { + args = append(args, "--cover-root", files.cover) + } + if err := runServer(context.Background(), args); err == nil || !strings.Contains(err.Error(), "--destination-write-timeout") { + t.Fatalf("%s check=%s accepted negative timeout: %v", mode, check, err) + } + } + } + if got := dnsCalls.Load(); got != 0 { + t.Fatalf("local timeout checks attempted %d DNS operations", got) + } +} + +func TestServerDestinationWriteTimeoutJSON(t *testing.T) { + clearPreflightEnvironment(t) + files := newPreflightFiles(t) + for _, mode := range []string{"native", "web"} { + for _, value := range []string{`"0s"`, `"37ms"`, `"-1s"`, `0`, `300`, `true`, `"bad-duration"`} { + t.Run(mode+"/"+value, func(t *testing.T) { + path := writeTestCommandConfig(t, `{"destination-write-timeout":`+value+`}`) + args := append(files.serverArgs(), "--protocol", mode, "--config", path) + if mode == "web" { + args = append(args, "--cover-root", files.cover) + } + err := runServer(context.Background(), args) + valid := value == `"0s"` || value == `"37ms"` + if (err == nil) != valid { + t.Fatalf("duration %s: error=%v valid=%t", value, err, valid) + } + }) + } + } +} + +func TestDestinationWriteTimeoutConfigPrecedence(t *testing.T) { + path := writeTestCommandConfig(t, `{"destination-write-timeout":"37s"}`) + for _, test := range []struct { + args []string + want time.Duration + }{ + {nil, 5 * time.Minute}, + {[]string{"--config", path}, 37 * time.Second}, + {[]string{"--config", path, "--destination-write-timeout=9s"}, 9 * time.Second}, + {[]string{"--config", path, "--destination-write-timeout=0s"}, 0}, + } { + fs := flag.NewFlagSet("server", flag.ContinueOnError) + fs.SetOutput(io.Discard) + value := fs.Duration("destination-write-timeout", 5*time.Minute, "") + if err := parseFlagsWithConfig(fs, test.args); err != nil { + t.Fatal(err) + } + if *value != test.want { + t.Fatalf("timeout=%s want=%s", *value, test.want) + } + } + for _, value := range []string{`0`, `"invalid"`} { + path := writeTestCommandConfig(t, `{"destination-write-timeout":`+value+`}`) + fs := flag.NewFlagSet("server", flag.ContinueOnError) + fs.SetOutput(io.Discard) + fs.Duration("destination-write-timeout", 5*time.Minute, "") + if err := parseFlagsWithConfig(fs, []string{"--config", path, "--destination-write-timeout=9s"}); err == nil { + t.Fatal("invalid config type/format was hidden by CLI override") + } + } +} + +func TestDestinationWriteTimeoutConfigIsServerOnly(t *testing.T) { + path := writeTestCommandConfig(t, `{"destination-write-timeout":"30s"}`) + if err := runClient(context.Background(), []string{"--config", path}); err == nil || !strings.Contains(err.Error(), "unknown or unsupported config option") { + t.Fatalf("client accepted the server-only destination timeout: %v", err) + } +} diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index c183130..487e1e6 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -253,6 +253,25 @@ progress as activity, including small buffered responses. Stalled request or response bodies and blocked writes remain bounded; header and keepalive timeouts are unchanged. +The relay's separate `--destination-write-timeout` defaults to `5m`; `0` +also selects `5m`, rather than disabling the limit. It bounds the completion +of each TCP destination write, split into chunks of at most 32 KiB. Each +successfully completed chunk gets a fresh budget for the next write. This is +not a kernel-level no-progress timer: a write can time out after making partial +progress, and that error is retained rather than silently retried. + +This server setting applies to native QUIC/TLS and web H2/H3 TCP tunnels, not +UDP or ordinary cover requests. The write deadline is cleared after every +write, including failures; when no write is pending it does not count idle +time, change read deadlines, or limit total upload duration. A stalled target +therefore releases its stream slot within the pending write's timeout, even +in the H3 response-FIN/reset edge case described in [WEB_COVER.md](WEB_COVER.md). +That is bounded cleanup, not a promise of immediate reset detection. Configure +it on the server as `--destination-write-timeout=30s` or JSON +`"destination-write-timeout": "30s"`; choose a longer value if individual +32 KiB writes can legitimately take longer. Negative durations are rejected by +normal startup and `--check`. + 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 fd57c81..f714cb2 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -45,11 +45,20 @@ 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. +FIN. After a clean response FIN, resetting only that H3 stream still cannot +directly interrupt an already blocked destination upload write through the +current QUIC API. The server's `--destination-write-timeout` now bounds that +pending write, releasing the stream slot without closing usable siblings. +Its default is `5m`; `0` selects that default and negative values are invalid. +This is bounded cleanup, not immediate reset detection. + +The limit is the completion deadline of each TCP destination write, in chunks +of at most 32 KiB—not a kernel-level no-progress timer or a total upload limit. +After a chunk completes successfully, the next write gets a new budget. +Partial-write errors remain errors. Deadlines are cleared after writes, do not run while no write +is pending, and never change destination read deadlines. Native QUIC/TLS TCP +tunnels share this setting; UDP and ordinary cover traffic do not. See +[deployment settings](DEPLOYMENT.md) for CLI and JSON examples. H2's handshake budget covers both TLS negotiation and the initial HTTP/2 preface/SETTINGS write. Caller cancellation or client shutdown also interrupts diff --git a/examples/server.json b/examples/server.json index 070abcf..d240ccf 100644 --- a/examples/server.json +++ b/examples/server.json @@ -3,5 +3,6 @@ "listen": ":8443", "cert": "server.crt", "key": "server.key", - "token-file": "relay-token" + "token-file": "relay-token", + "destination-write-timeout": "5m" } diff --git a/internal/tunnel/common.go b/internal/tunnel/common.go index 18a92ab..2f4fc74 100644 --- a/internal/tunnel/common.go +++ b/internal/tunnel/common.go @@ -21,11 +21,12 @@ import ( ) const ( - defaultHandshakeTimeout = 10 * time.Second - defaultDialTimeout = 10 * time.Second - defaultMaxStreams = 1024 - defaultMaxConnections = 256 - defaultMaxClientConnections = 32 + defaultHandshakeTimeout = 10 * time.Second + defaultDialTimeout = 10 * time.Second + defaultDestinationWriteTimeout = 5 * time.Minute + defaultMaxStreams = 1024 + defaultMaxConnections = 256 + defaultMaxClientConnections = 32 ) // RemoteError is returned when the authenticated exit rejects a CONNECT @@ -44,11 +45,12 @@ func (e *RemoteError) Error() string { } type serverCore struct { - tokenHash [sha256.Size]byte - dialer transport.Dialer - handshakeTimeout time.Duration - dialTimeout time.Duration - sem chan struct{} + tokenHash [sha256.Size]byte + dialer transport.Dialer + handshakeTimeout time.Duration + dialTimeout time.Duration + destinationWriteTimeout time.Duration + sem chan struct{} } // StreamAdmission is a concurrency budget for active relay streams. Pass the @@ -70,7 +72,7 @@ func NewStreamAdmission(limit int) (*StreamAdmission, error) { } func newServerCore(token string, dialer transport.Dialer, handshakeTimeout, dialTimeout time.Duration, maxStreams int) (*serverCore, error) { - return newServerCoreWithAdmission(token, dialer, handshakeTimeout, dialTimeout, maxStreams, nil) + return newServerCoreWithAdmission(token, dialer, handshakeTimeout, dialTimeout, 0, maxStreams, nil) } func newServerCoreWithAdmission( @@ -78,6 +80,7 @@ func newServerCoreWithAdmission( dialer transport.Dialer, handshakeTimeout time.Duration, dialTimeout time.Duration, + destinationWriteTimeout time.Duration, maxStreams int, admission *StreamAdmission, ) (*serverCore, error) { @@ -88,7 +91,7 @@ func newServerCoreWithAdmission( netDialer := &net.Dialer{Timeout: defaultDialTimeout, KeepAlive: 30 * time.Second} dialer = netDialer } - if handshakeTimeout < 0 || dialTimeout < 0 || maxStreams < 0 { + if handshakeTimeout < 0 || dialTimeout < 0 || destinationWriteTimeout < 0 || maxStreams < 0 { return nil, errors.New("tunnel: timeout and concurrency limits cannot be negative") } if handshakeTimeout == 0 { @@ -97,6 +100,9 @@ func newServerCoreWithAdmission( if dialTimeout == 0 { dialTimeout = defaultDialTimeout } + if destinationWriteTimeout == 0 { + destinationWriteTimeout = defaultDestinationWriteTimeout + } if maxStreams == 0 { if admission == nil { maxStreams = defaultMaxStreams @@ -122,11 +128,12 @@ func newServerCoreWithAdmission( } } return &serverCore{ - tokenHash: sha256.Sum256([]byte(token)), - dialer: dialer, - handshakeTimeout: handshakeTimeout, - dialTimeout: dialTimeout, - sem: admission.sem, + tokenHash: sha256.Sum256([]byte(token)), + dialer: dialer, + handshakeTimeout: handshakeTimeout, + dialTimeout: dialTimeout, + destinationWriteTimeout: destinationWriteTimeout, + sem: admission.sem, }, nil } @@ -217,6 +224,7 @@ func (s *serverCore) handleStream( _ = protocol.WriteResponse(stream, protocol.Response{Status: protocol.StatusDialFailed, Message: "destination unavailable"}) return } + upstream = s.boundDestinationWrites(upstream) defer upstream.Close() options.Response.Status = protocol.StatusOK if err := protocol.WriteResponse(stream, options.Response); err != nil { diff --git a/internal/tunnel/destination_write.go b/internal/tunnel/destination_write.go new file mode 100644 index 0000000..5271b9b --- /dev/null +++ b/internal/tunnel/destination_write.go @@ -0,0 +1,72 @@ +package tunnel + +import ( + "errors" + "io" + "net" + "sync" + "time" +) + +const destinationWriteChunkSize = 32 * 1024 + +func (s *serverCore) boundDestinationWrites(conn net.Conn) net.Conn { + return &destinationWriteConn{Conn: conn, timeout: s.destinationWriteTimeout} +} + +// destinationWriteConn owns the write deadline of a newly dialed TCP target. +// A pending write has a completion budget, not a kernel-level inactivity timer: +// partial progress inside net.Conn.Write cannot be observed until it returns. +// No deadline remains armed while there is no pending write, and read deadlines +// are untouched. Embed net.Conn rather than its concrete implementation so +// io.Copy cannot bypass Write via a promoted ReaderFrom or WriterTo method. +type destinationWriteConn struct { + net.Conn + timeout time.Duration + writeMu sync.Mutex +} + +func (c *destinationWriteConn) Write(p []byte) (int, error) { + c.writeMu.Lock() + defer c.writeMu.Unlock() + written := 0 + for len(p) > 0 { + chunk := p[:min(len(p), destinationWriteChunkSize)] + if err := c.Conn.SetWriteDeadline(time.Now().Add(c.timeout)); err != nil { + return written, err + } + n, err := c.Conn.Write(chunk) + clearErr := c.Conn.SetWriteDeadline(time.Time{}) + if n < 0 || n > len(chunk) { + return written, errors.New("tunnel: invalid destination write count") + } + written += n + if err != nil { + // Preserve partial progress and its error; retrying a timed-out + // write here could hide a failed destination from the relay. + return written, err + } + if clearErr != nil { + return written, clearErr + } + if n == 0 { + return written, io.ErrShortWrite + } + p = p[n:] + } + return written, nil +} + +func (c *destinationWriteConn) CloseWrite() error { + if conn, ok := c.Conn.(closeWriter); ok { + return conn.CloseWrite() + } + return nil +} + +func (c *destinationWriteConn) CloseRead() error { + if conn, ok := c.Conn.(closeReader); ok { + return conn.CloseRead() + } + return nil +} diff --git a/internal/tunnel/destination_write_integration_test.go b/internal/tunnel/destination_write_integration_test.go new file mode 100644 index 0000000..4c620f4 --- /dev/null +++ b/internal/tunnel/destination_write_integration_test.go @@ -0,0 +1,302 @@ +package tunnel + +import ( + "bytes" + "context" + "fmt" + "io" + "net" + "net/http" + "sync" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +const destinationIntegrationWriteTimeout = 200 * time.Millisecond + +// These are real authenticated client/server pairs. Only the destination is +// a controlled net.Pipe, whose unread peer guarantees a blocked network Write. +func TestDestinationWriteTimeoutReleasesStalledRelay(t *testing.T) { + for _, mode := range []string{"quic", "tls", "h2", "h3"} { + t.Run(mode, func(t *testing.T) { + exerciseDestinationWriteTimeout(t, mode, false) + }) + } +} + +func TestH3DestinationWriteTimeoutAfterResponseFINAndReset(t *testing.T) { + exerciseDestinationWriteTimeout(t, "h3", true) +} + +func exerciseDestinationWriteTimeout(t *testing.T, mode string, resetAfterFIN bool) { + t.Helper() + upstream, target := net.Pipe() + defer target.Close() + tracked := &destinationIntegrationTarget{ + Conn: upstream, started: make(chan time.Time, 1), result: make(chan error, 1), closed: make(chan struct{}), + } + defer tracked.Close() // Also releases the baseline failure before server cleanup. + 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 resetAfterFIN { + return &destinationIntegrationResponseFIN{tracked}, nil + } + return tracked, nil + } + return (&net.Dialer{}).DialContext(ctx, network, address) + }) + client := newDestinationWriteIntegrationClient(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 resetAfterFIN { + if response, err := io.ReadAll(conn); err != nil || len(response) != 0 { + t.Fatalf("response FIN = %q, %v", response, err) + } + } + if _, err := io.WriteString(conn, "upload"); err != nil { + t.Fatal(err) + } + var started time.Time + select { + case started = <-tracked.started: + case <-time.After(time.Second): + t.Fatal("upload never entered the destination Write") + } + if resetAfterFIN { + _ = conn.Close() + } + // No unrelated cancellation is used for the ordinary four-transport + // cases. The post-FIN reset case must likewise be a write timeout, not + // an assertion that the currently unavailable reset signal is immediate. + select { + case err := <-tracked.result: + if timeout, ok := err.(net.Error); !ok || !timeout.Timeout() { + t.Fatalf("destination Write error = %v, want its configured timeout", err) + } + if elapsed := time.Since(started); elapsed < destinationIntegrationWriteTimeout/2 { + t.Fatalf("Write ended after %s, too early to prove bounded stall reclamation", elapsed) + } + case <-time.After(time.Second): + t.Fatal("stalled destination Write outlived its 200 ms timeout") + } + select { + case <-tracked.closed: + case <-time.After(time.Second): + t.Fatal("expired destination remained open") + } + awaitWebServerAdmission(t, admission, 1) + assertWebSessionSiblingEcho(t, sibling) + _ = sibling.Close() + awaitWebServerAdmission(t, admission, 0) +} + +// A write deadline must be scoped to a pending Write, not a dormant half-open +// connection or the entire upload. This target reads normally after its FIN. +func TestH3DestinationWriteTimeoutPreservesResponseFINUpload(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() + chunk := []byte("after-FIN-upload\x00\xff") + payload := bytes.Repeat(chunk, 8) + 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, payload) { + return fmt.Errorf("upload = %q, want %q", got, payload) + } + return nil + }() + }() + t.Cleanup(func() { + cancel() + _ = listener.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Error("response-FIN target did not stop") + } + }) + admission, err := NewStreamAdmission(1) + if err != nil { + t.Fatal(err) + } + client := newDestinationWriteIntegrationClient(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) + } + response, err := io.ReadAll(conn) + if err != nil || !bytes.Equal(response, wantResponse) { + t.Fatalf("response FIN = %q, %v", response, err) + } + // With no pending Write, waiting longer than the bound is harmless. + time.Sleep(2 * destinationIntegrationWriteTimeout) + for range 8 { + if _, err := conn.Write(chunk); err != nil { + t.Fatalf("active upload after response FIN: %v", err) + } + time.Sleep(50 * time.Millisecond) + } + 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("normal upload did not finish after response FIN") + } + awaitWebServerAdmission(t, admission, 0) +} + +type destinationIntegrationServer interface { + Addr() net.Addr + Serve(context.Context) error + Close() error +} + +func newDestinationWriteIntegrationClient(t *testing.T, mode string, outbound transport.Dialer, admission *StreamAdmission) transport.Dialer { + t.Helper() + serverTLS, clientTLS := testTLSConfigs(t) + var server destinationIntegrationServer + var err error + switch mode { + case "quic": + server, err = ListenQUIC(QUICServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + Dialer: outbound, StreamAdmission: admission, DestinationWriteTimeout: destinationIntegrationWriteTimeout, + }) + case "tls": + server, err = ListenTLS(TLSServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + Dialer: outbound, StreamAdmission: admission, DestinationWriteTimeout: destinationIntegrationWriteTimeout, + }) + case "h2": + server, err = ListenWebH2(WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, Cover: http.NotFoundHandler(), + Dialer: outbound, StreamAdmission: admission, DestinationWriteTimeout: destinationIntegrationWriteTimeout, + }) + case "h3": + server, err = ListenWebH3(WebH3ServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, Cover: http.NotFoundHandler(), + Dialer: outbound, StreamAdmission: admission, DestinationWriteTimeout: destinationIntegrationWriteTimeout, + }) + default: + t.Fatalf("unknown transport %q", mode) + } + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- server.Serve(ctx) }() + t.Cleanup(func() { + cancel() + _ = server.Close() + select { + case err := <-done: + if err != nil { + t.Errorf("%s server: %v", mode, err) + } + case <-time.After(time.Second): + t.Errorf("%s server did not stop", mode) + } + }) + var client interface { + transport.Dialer + io.Closer + } + switch mode { + case "quic": + client, err = NewClient(ClientConfig{ServerAddress: server.Addr().String(), Token: testToken, TLSConfig: clientTLS}) + case "tls": + client, err = NewTLSClient(TLSClientConfig{ServerAddress: server.Addr().String(), Token: testToken, TLSConfig: clientTLS}) + case "h2": + client, err = NewWebH2Client(WebH2ClientConfig{ServerAddress: server.Addr().String(), Token: testToken, TLSConfig: clientTLS}) + case "h3": + client, err = NewWebH3Client(WebH3ClientConfig{ServerAddress: server.Addr().String(), Token: testToken, TLSConfig: clientTLS}) + } + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + return client +} + +type destinationIntegrationTarget struct { + net.Conn + started chan time.Time + result chan error + closed chan struct{} + writeOnce, closeOnce sync.Once +} + +func (c *destinationIntegrationTarget) Write(p []byte) (int, error) { + c.writeOnce.Do(func() { c.started <- time.Now() }) + n, err := c.Conn.Write(p) + c.result <- err + return n, err +} + +func (c *destinationIntegrationTarget) Close() error { + var err error + c.closeOnce.Do(func() { err = c.Conn.Close(); close(c.closed) }) + return err +} + +type destinationIntegrationResponseFIN struct{ *destinationIntegrationTarget } + +func (*destinationIntegrationResponseFIN) Read([]byte) (int, error) { return 0, io.EOF } diff --git a/internal/tunnel/destination_write_test.go b/internal/tunnel/destination_write_test.go new file mode 100644 index 0000000..c4fd592 --- /dev/null +++ b/internal/tunnel/destination_write_test.go @@ -0,0 +1,276 @@ +package tunnel + +import ( + "bytes" + "errors" + "io" + "net" + "sync" + "testing" + "time" +) + +func TestDestinationWriteChunksCannotBeBypassedByCopy(t *testing.T) { + for _, writerTo := range []bool{false, true} { + name := "reader_only" + if writerTo { + name = "source_writer_to" + } + t.Run(name, func(t *testing.T) { + payload := bytes.Repeat([]byte("x"), 3*destinationWriteChunkSize+1) + raw := &destinationWriteSpy{} + conn := &destinationWriteConn{Conn: raw, timeout: time.Second} + if _, ok := any(conn).(io.ReaderFrom); ok { + t.Fatal("io.Copy could bypass deadline handling through ReaderFrom") + } + if _, ok := any(conn).(io.WriterTo); ok { + t.Fatal("io.Copy could bypass the wrapper through WriterTo") + } + var source io.Reader = bytes.NewReader(payload) + if !writerTo { + source = struct{ io.Reader }{source} + } + before := time.Now() + n, err := io.Copy(conn, source) + if err != nil || n != int64(len(payload)) || !bytes.Equal(raw.data.Bytes(), payload) { + t.Fatalf("copy = %d, %v, target bytes=%d", n, err, raw.data.Len()) + } + if raw.readFromCalled || raw.maxWrite > destinationWriteChunkSize { + t.Fatalf("bypassed chunking: ReadFrom=%v, largest write=%d", raw.readFromCalled, raw.maxWrite) + } + if raw.writes != 4 || len(raw.deadlines) != 8 { + t.Fatalf("writes/deadline changes = %d/%d, want 4/8", raw.writes, len(raw.deadlines)) + } + for i := 0; i < len(raw.deadlines); i += 2 { + if raw.deadlines[i].Before(before.Add(time.Second)) || !raw.deadlines[i+1].IsZero() { + t.Fatalf("write %d did not receive and then clear its deadline", i/2) + } + } + if raw.readDeadlines != 0 { + t.Fatal("destination write policy changed a read deadline") + } + if err := conn.CloseWrite(); err != nil || raw.halfCloses != 1 { + t.Fatal("destination half-close was not forwarded") + } + }) + } +} + +func TestDestinationWriteErrorsPreservePartialProgress(t *testing.T) { + writeErr := errors.New("destination write failed") + deadlineErr := errors.New("deadline update failed") + for _, tc := range []struct { + name string + n int + err error + failDeadline int + wantN int + wantErr error + wantWrites int + }{ + {"partial_error", 2, writeErr, 0, 2, writeErr, 1}, + {"partial_timeout", 2, &net.DNSError{IsTimeout: true}, 0, 2, nil, 1}, + {"zero_progress", 0, nil, 0, 0, io.ErrShortWrite, 1}, + {"arm_error", 0, nil, 1, 0, deadlineErr, 0}, + {"clear_error", 4, nil, 2, 4, deadlineErr, 1}, + {"write_error_wins", 2, writeErr, 2, 2, writeErr, 1}, + } { + t.Run(tc.name, func(t *testing.T) { + raw := &destinationWriteSpy{ + write: func([]byte) (int, error) { return tc.n, tc.err }, + failDeadline: tc.failDeadline, deadlineErr: deadlineErr, + } + conn := &destinationWriteConn{Conn: raw, timeout: time.Second} + n, err := conn.Write([]byte("data")) + wantErr := tc.wantErr + if tc.name == "partial_timeout" { + wantErr = tc.err + } + if n != tc.wantN || !errors.Is(err, wantErr) || raw.writes != tc.wantWrites { + t.Fatalf("Write=%d,%v writes=%d; want %d,%v writes=%d", n, err, raw.writes, tc.wantN, wantErr, tc.wantWrites) + } + if tc.failDeadline != 1 && (len(raw.deadlines) != 2 || !raw.deadlines[1].IsZero()) { + t.Fatal("completed or failed write did not clear its deadline") + } + }) + } +} + +func TestDestinationWriteRejectsInvalidCounts(t *testing.T) { + for _, count := range []int{-1, 5} { + raw := &destinationWriteSpy{write: func([]byte) (int, error) { return count, nil }} + conn := &destinationWriteConn{Conn: raw, timeout: time.Second} + if n, err := conn.Write([]byte("data")); n != 0 || err == nil { + t.Fatalf("invalid count %d returned %d, %v", count, n, err) + } + if len(raw.deadlines) != 2 || !raw.deadlines[1].IsZero() { + t.Fatal("invalid count left its deadline armed") + } + } +} + +func TestDestinationWriteContinuesSuccessfulShortWrites(t *testing.T) { + raw := &destinationWriteSpy{} + raw.write = func(p []byte) (int, error) { + return raw.data.Write(p[:min(len(p), 2)]) + } + conn := &destinationWriteConn{Conn: raw, timeout: time.Second} + payload := []byte("payload") + if n, err := conn.Write(payload); n != len(payload) || err != nil || !bytes.Equal(raw.data.Bytes(), payload) { + t.Fatalf("short-write continuation = %d, %v, %q", n, err, raw.data.Bytes()) + } + if raw.writes != 4 || len(raw.deadlines) != 8 { + t.Fatalf("short writes/deadline updates = %d/%d, want 4/8", raw.writes, len(raw.deadlines)) + } +} + +func TestDestinationWriteSuccessLeavesIdleAndReadsAlone(t *testing.T) { + local, peer := net.Pipe() + defer local.Close() + defer peer.Close() + if err := local.SetDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatal(err) + } + if err := peer.SetDeadline(time.Now().Add(2 * time.Second)); err != nil { + t.Fatal(err) + } + conn := &destinationWriteConn{Conn: local, timeout: 80 * time.Millisecond} + done := make(chan error, 1) + go func() { + var payload [5]byte + if _, err := io.ReadFull(peer, payload[:]); err != nil { + done <- err + return + } + if _, err := peer.Write([]byte("reply")); err != nil { + done <- err + return + } + _, err := io.ReadFull(peer, payload[:]) + done <- err + }() + if _, err := conn.Write([]byte("first")); err != nil { + t.Fatal(err) + } + time.Sleep(2 * conn.timeout) + var reply [5]byte + if _, err := io.ReadFull(conn, reply[:]); err != nil || string(reply[:]) != "reply" { + t.Fatalf("read after successful write and idle = %q, %v", reply, err) + } + if _, err := conn.Write([]byte("later")); err != nil { + t.Fatalf("write after idle: %v", err) + } + select { + case err := <-done: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("target did not finish after resumed write") + } +} + +func TestDestinationWriteDeadlineBoundsPendingIO(t *testing.T) { + local, peer := net.Pipe() + defer local.Close() + defer peer.Close() + conn := &destinationWriteConn{Conn: local, timeout: 80 * time.Millisecond} + done := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("blocked upload")) + done <- err + }() + select { + case err := <-done: + var timeout net.Error + if !errors.As(err, &timeout) || !timeout.Timeout() { + t.Fatalf("write error = %v, want timeout", err) + } + case <-time.After(time.Second): + t.Fatal("destination write was not bounded") + } +} + +func TestDestinationWriteCloseDoesNotWaitForWriter(t *testing.T) { + local, peer := net.Pipe() + defer local.Close() + defer peer.Close() + entered := make(chan struct{}) + raw := &destinationWriteSignalConn{Conn: local, entered: entered} + conn := &destinationWriteConn{Conn: raw, timeout: time.Hour} + done := make(chan error, 2) + for range 2 { + go func() { + _, err := conn.Write([]byte("blocked")) + done <- err + }() + } + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("writer did not enter target Write") + } + closed := make(chan struct{}) + go func() { _ = conn.Close(); close(closed) }() + select { + case <-closed: + case <-time.After(time.Second): + t.Fatal("Close waited for the blocked writer lock") + } + for range 2 { + select { + case err := <-done: + if err == nil { + t.Fatal("blocked write succeeded after Close") + } + case <-time.After(time.Second): + t.Fatal("writer remained blocked after Close") + } + } +} + +type destinationWriteSignalConn struct { + net.Conn + entered chan struct{} + once sync.Once +} + +func (c *destinationWriteSignalConn) Write(p []byte) (int, error) { + c.once.Do(func() { close(c.entered) }) + return c.Conn.Write(p) +} + +type destinationWriteSpy struct { + net.Conn + data bytes.Buffer + writes, maxWrite, readDeadlines, halfCloses int + deadlines []time.Time + readFromCalled bool + write func([]byte) (int, error) + failDeadline int + deadlineErr error +} + +func (c *destinationWriteSpy) Write(p []byte) (int, error) { + c.writes++ + c.maxWrite = max(c.maxWrite, len(p)) + if c.write != nil { + return c.write(p) + } + return c.data.Write(p) +} + +func (c *destinationWriteSpy) SetWriteDeadline(deadline time.Time) error { + c.deadlines = append(c.deadlines, deadline) + if len(c.deadlines) == c.failDeadline { + return c.deadlineErr + } + return nil +} + +func (c *destinationWriteSpy) SetReadDeadline(time.Time) error { c.readDeadlines++; return nil } +func (c *destinationWriteSpy) CloseWrite() error { c.halfCloses++; return nil } +func (c *destinationWriteSpy) ReadFrom(io.Reader) (int64, error) { + c.readFromCalled = true + return 0, errors.New("unbounded ReaderFrom invoked") +} diff --git a/internal/tunnel/quic.go b/internal/tunnel/quic.go index ee98ec3..9636aeb 100644 --- a/internal/tunnel/quic.go +++ b/internal/tunnel/quic.go @@ -26,14 +26,17 @@ const ( // QUICServerConfig configures an encrypted QUIC exit listener. type QUICServerConfig struct { - Address string - Token string - TLSConfig *tls.Config - QUICConfig *quic.Config - Dialer transport.Dialer - HandshakeTimeout time.Duration - DialTimeout time.Duration - MaxConcurrentStreams int + Address string + Token string + TLSConfig *tls.Config + QUICConfig *quic.Config + Dialer transport.Dialer + HandshakeTimeout time.Duration + DialTimeout time.Duration + // DestinationWriteTimeout bounds each TCP destination write. Zero uses + // five minutes; negative values are invalid. It does not apply to UDP. + DestinationWriteTimeout time.Duration + MaxConcurrentStreams int // StreamAdmission optionally shares the active-stream budget with other // server transports. When set, MaxConcurrentStreams must be zero or equal // to the admission limit. Nil preserves the independent-server behavior. @@ -116,6 +119,7 @@ func ListenQUIC(config QUICServerConfig) (*QUICServer, error) { config.Dialer, config.HandshakeTimeout, config.DialTimeout, + config.DestinationWriteTimeout, config.MaxConcurrentStreams, config.StreamAdmission, ) diff --git a/internal/tunnel/server_destination_timeout_test.go b/internal/tunnel/server_destination_timeout_test.go new file mode 100644 index 0000000..1ab090b --- /dev/null +++ b/internal/tunnel/server_destination_timeout_test.go @@ -0,0 +1,95 @@ +package tunnel + +import ( + "io" + "net" + "net/http" + "testing" + "time" +) + +func TestServerDestinationWriteTimeoutConfigurationWiring(t *testing.T) { + serverTLS, _ := testTLSConfigs(t) + for _, mode := range []string{"tls", "quic", "h2", "h3", "web"} { + for _, value := range []time.Duration{-time.Second, 0, 37 * time.Millisecond} { + t.Run(mode+"/"+value.String(), func(t *testing.T) { + var ( + core *serverCore + closer io.Closer + err error + ) + switch mode { + case "tls": + server, createErr := ListenTLS(TLSServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + DestinationWriteTimeout: value, + }) + err = createErr + if err == nil { + core, closer = server.core, server + } + case "quic": + server, createErr := ListenQUIC(QUICServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + DestinationWriteTimeout: value, + }) + err = createErr + if err == nil { + core, closer = server.core, server + } + case "h2": + server, createErr := ListenWebH2(WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + Cover: http.NotFoundHandler(), Dialer: &net.Dialer{}, DestinationWriteTimeout: value, + }) + err = createErr + if err == nil { + core, closer = server.server.Handler.(*webTunnelHandler).core, server + } + case "h3": + server, createErr := ListenWebH3(WebH3ServerConfig{ + Address: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + Cover: http.NotFoundHandler(), Dialer: &net.Dialer{}, DestinationWriteTimeout: value, + }) + err = createErr + if err == nil { + core, closer = server.server.Handler.(*webTunnelHandler).core, server + } + case "web": + server, createErr := ListenWeb(WebServerConfig{ + TCPAddress: "127.0.0.1:0", Token: testToken, TLSConfig: serverTLS, + Cover: http.NotFoundHandler(), Dialer: &net.Dialer{}, DestinationWriteTimeout: value, + }) + err = createErr + if err == nil { + core, closer = server.h2.server.Handler.(*webTunnelHandler).core, server + h3Core := server.h3.server.Handler.(*webTunnelHandler).core + if h3Core != core { + _ = server.Close() + t.Fatal("combined web server did not share the timeout-bearing core") + } + } + } + if closer != nil { + t.Cleanup(func() { _ = closer.Close() }) + } + if value < 0 { + if err == nil { + t.Fatal("negative destination write timeout was accepted") + } + return + } + if err != nil { + t.Fatal(err) + } + want := value + if want == 0 { + want = 5 * time.Minute + } + if core.destinationWriteTimeout != want { + t.Fatalf("destination write timeout=%s want=%s", core.destinationWriteTimeout, want) + } + }) + } + } +} diff --git a/internal/tunnel/stream_admission_test.go b/internal/tunnel/stream_admission_test.go index 8097fbc..bef9e16 100644 --- a/internal/tunnel/stream_admission_test.go +++ b/internal/tunnel/stream_admission_test.go @@ -24,11 +24,11 @@ func TestStreamAdmissionValidationAndCompatibility(t *testing.T) { if err != nil { t.Fatal(err) } - first, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 1, shared) + first, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 0, 1, shared) if err != nil { t.Fatal(err) } - second, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 0, shared) + second, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 0, 0, shared) if err != nil { t.Fatal(err) } @@ -70,10 +70,10 @@ func TestStreamAdmissionRejectsInvalidConfiguration(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 1, shared); err == nil { + if _, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 0, 1, shared); err == nil { t.Fatal("mismatched MaxConcurrentStreams and StreamAdmission were accepted") } - if _, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 0, &StreamAdmission{}); err == nil { + if _, err := newServerCoreWithAdmission(testToken, nil, 0, 0, 0, 0, &StreamAdmission{}); err == nil { t.Fatal("zero-value StreamAdmission was accepted") } } diff --git a/internal/tunnel/tls.go b/internal/tunnel/tls.go index 2ec3d29..151f75c 100644 --- a/internal/tunnel/tls.go +++ b/internal/tunnel/tls.go @@ -17,13 +17,16 @@ import ( // TLSServerConfig configures the TCP+TLS fallback exit listener. type TLSServerConfig struct { - Address string - Token string - TLSConfig *tls.Config - Dialer transport.Dialer - HandshakeTimeout time.Duration - DialTimeout time.Duration - MaxConcurrentStreams int + Address string + Token string + TLSConfig *tls.Config + Dialer transport.Dialer + HandshakeTimeout time.Duration + DialTimeout time.Duration + // DestinationWriteTimeout bounds each TCP destination write. Zero uses + // five minutes; negative values are invalid. Idle reads are unaffected. + DestinationWriteTimeout time.Duration + MaxConcurrentStreams int // StreamAdmission optionally shares the active-stream budget with other // server transports. When set, MaxConcurrentStreams must be zero or equal // to the admission limit. Nil preserves the independent-server behavior. @@ -68,6 +71,7 @@ func ListenTLS(config TLSServerConfig) (*TLSServer, error) { config.Dialer, config.HandshakeTimeout, config.DialTimeout, + config.DestinationWriteTimeout, config.MaxConcurrentStreams, config.StreamAdmission, ) diff --git a/internal/tunnel/web_h2.go b/internal/tunnel/web_h2.go index d090cc0..fdc9bd5 100644 --- a/internal/tunnel/web_h2.go +++ b/internal/tunnel/web_h2.go @@ -29,15 +29,18 @@ type webTLSConnectionContextKey struct{} // Only a TLS 1.3 authenticated HTTP/2 CONNECT request is handled as a tunnel; // all other requests are delegated to Cover. type WebH2ServerConfig struct { - Address string - Token string - TLSConfig *tls.Config - Cover http.Handler - Dialer transport.Dialer - HandshakeTimeout time.Duration - DialTimeout time.Duration - MaxConcurrentStreams int - StreamAdmission *StreamAdmission + Address string + Token string + TLSConfig *tls.Config + Cover http.Handler + Dialer transport.Dialer + HandshakeTimeout time.Duration + DialTimeout time.Duration + // DestinationWriteTimeout bounds each tunneled TCP destination write. + // Zero uses five minutes; negative values are invalid. Cover is unaffected. + DestinationWriteTimeout time.Duration + MaxConcurrentStreams int + StreamAdmission *StreamAdmission // MaxConnections and MaxClientConnections bound accepted HTTPS // connections globally and per source IPv4 or IPv6 /64. A combined // WebServer shares these limits with HTTP/3. @@ -74,6 +77,7 @@ func ListenWebH2(config WebH2ServerConfig) (*WebH2Server, error) { config.Dialer, config.HandshakeTimeout, config.DialTimeout, + config.DestinationWriteTimeout, config.MaxConcurrentStreams, config.StreamAdmission, ) diff --git a/internal/tunnel/web_h3.go b/internal/tunnel/web_h3.go index aff6fb5..7dd95ea 100644 --- a/internal/tunnel/web_h3.go +++ b/internal/tunnel/web_h3.go @@ -18,16 +18,19 @@ import ( // WebH3ServerConfig configures a real HTTP/3 cover origin whose authenticated // CONNECT requests carry TCP proxy streams. type WebH3ServerConfig struct { - Address string - Token string - TLSConfig *tls.Config - QUICConfig *quic.Config - Dialer transport.Dialer - Cover http.Handler - HandshakeTimeout time.Duration - DialTimeout time.Duration - MaxConcurrentStreams int - StreamAdmission *StreamAdmission + Address string + Token string + TLSConfig *tls.Config + QUICConfig *quic.Config + Dialer transport.Dialer + Cover http.Handler + HandshakeTimeout time.Duration + DialTimeout time.Duration + // DestinationWriteTimeout bounds each tunneled TCP destination write. + // Zero uses five minutes; negative values are invalid. Cover/UDP are unaffected. + DestinationWriteTimeout time.Duration + MaxConcurrentStreams int + StreamAdmission *StreamAdmission // MaxConnections and MaxClientConnections bound accepted HTTPS // connections globally and per source IPv4 or IPv6 /64. A combined // WebServer shares these limits with HTTP/2. @@ -83,6 +86,7 @@ func ListenWebH3(config WebH3ServerConfig) (*WebH3Server, error) { config.Dialer, config.HandshakeTimeout, config.DialTimeout, + config.DestinationWriteTimeout, config.MaxConcurrentStreams, config.StreamAdmission, ) diff --git a/internal/tunnel/web_handler.go b/internal/tunnel/web_handler.go index a68dc04..f2293ad 100644 --- a/internal/tunnel/web_handler.go +++ b/internal/tunnel/web_handler.go @@ -69,6 +69,7 @@ func (h *webTunnelHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { h.writeAuthenticatedError(w, &authentication, http.StatusBadGateway) return } + upstream = h.core.boundDestinationWrites(upstream) defer upstream.Close() switch wire { diff --git a/internal/tunnel/web_server.go b/internal/tunnel/web_server.go index 125121a..5eecadf 100644 --- a/internal/tunnel/web_server.go +++ b/internal/tunnel/web_server.go @@ -32,10 +32,13 @@ type WebServerConfig struct { Cover http.Handler Dialer transport.Dialer - HandshakeTimeout time.Duration - DialTimeout time.Duration - MaxConcurrentStreams int - StreamAdmission *StreamAdmission + HandshakeTimeout time.Duration + DialTimeout time.Duration + // DestinationWriteTimeout bounds each tunneled TCP destination write. + // Zero uses five minutes; negative values are invalid. Cover/UDP are unaffected. + DestinationWriteTimeout time.Duration + MaxConcurrentStreams int + StreamAdmission *StreamAdmission // MaxConnections and MaxClientConnections are shared by the HTTP/2 and // HTTP/3 listeners, preventing a source from multiplying its allowance by // switching transports. @@ -86,6 +89,7 @@ func ListenWeb(config WebServerConfig) (*WebServer, error) { config.Dialer, config.HandshakeTimeout, config.DialTimeout, + config.DestinationWriteTimeout, config.MaxConcurrentStreams, config.StreamAdmission, ) @@ -111,20 +115,21 @@ func ListenWeb(config WebServerConfig) (*WebServer, error) { // unrelated alternative service. tcpCover := &webAltSvcCover{next: config.Cover} h2, err := listenWebH2WithCore(WebH2ServerConfig{ - Address: config.TCPAddress, - Token: config.Token, - TLSConfig: config.TLSConfig, - Cover: tcpCover, - Dialer: config.Dialer, - HandshakeTimeout: config.HandshakeTimeout, - DialTimeout: config.DialTimeout, - MaxConcurrentStreams: config.MaxConcurrentStreams, - StreamAdmission: config.StreamAdmission, - ReplayEntries: config.ReplayEntries, - MaxConnections: config.MaxConnections, - MaxClientConnections: config.MaxClientConnections, - MaxHeaderBytes: config.MaxHeaderBytes, - connectionAdmission: connectionAdmission, + Address: config.TCPAddress, + Token: config.Token, + TLSConfig: config.TLSConfig, + Cover: tcpCover, + Dialer: config.Dialer, + HandshakeTimeout: config.HandshakeTimeout, + DialTimeout: config.DialTimeout, + DestinationWriteTimeout: config.DestinationWriteTimeout, + MaxConcurrentStreams: config.MaxConcurrentStreams, + StreamAdmission: config.StreamAdmission, + ReplayEntries: config.ReplayEntries, + MaxConnections: config.MaxConnections, + MaxClientConnections: config.MaxClientConnections, + MaxHeaderBytes: config.MaxHeaderBytes, + connectionAdmission: connectionAdmission, }, core, verifier) if err != nil { return nil, err @@ -136,26 +141,27 @@ func ListenWeb(config WebServerConfig) (*WebServer, error) { return nil, err } h3, err := listenWebH3WithCore(WebH3ServerConfig{ - Address: udpAddress, - Token: config.Token, - TLSConfig: config.TLSConfig, - QUICConfig: config.QUICConfig, - Dialer: config.Dialer, - Cover: config.Cover, - HandshakeTimeout: config.HandshakeTimeout, - DialTimeout: config.DialTimeout, - MaxConcurrentStreams: config.MaxConcurrentStreams, - StreamAdmission: config.StreamAdmission, - MaxHeaderBytes: config.MaxHeaderBytes, - ReplayEntries: config.ReplayEntries, - MaxConnections: config.MaxConnections, - MaxClientConnections: config.MaxClientConnections, - connectionAdmission: connectionAdmission, - UDPResolver: config.UDPResolver, - MaxUDPSessions: config.MaxUDPSessions, - MaxClientUDPSessions: config.MaxClientUDPSessions, - MaxUDPDestinations: config.MaxUDPDestinations, - UDPReceiveQueue: config.UDPReceiveQueue, + Address: udpAddress, + Token: config.Token, + TLSConfig: config.TLSConfig, + QUICConfig: config.QUICConfig, + Dialer: config.Dialer, + Cover: config.Cover, + HandshakeTimeout: config.HandshakeTimeout, + DialTimeout: config.DialTimeout, + DestinationWriteTimeout: config.DestinationWriteTimeout, + MaxConcurrentStreams: config.MaxConcurrentStreams, + StreamAdmission: config.StreamAdmission, + MaxHeaderBytes: config.MaxHeaderBytes, + ReplayEntries: config.ReplayEntries, + MaxConnections: config.MaxConnections, + MaxClientConnections: config.MaxClientConnections, + connectionAdmission: connectionAdmission, + UDPResolver: config.UDPResolver, + MaxUDPSessions: config.MaxUDPSessions, + MaxClientUDPSessions: config.MaxClientUDPSessions, + MaxUDPDestinations: config.MaxUDPDestinations, + UDPReceiveQueue: config.UDPReceiveQueue, }, core, verifier) if err != nil { closeErr := normalizeWebServerCloseError(h2.Close())