diff --git a/internal/home/client.go b/internal/home/client.go index dc299b8c7..040805019 100644 --- a/internal/home/client.go +++ b/internal/home/client.go @@ -582,9 +582,9 @@ func (c *Client) trackedRedisDialer(dialer func(context.Context, string, string) } } -func (c *homeDispatchConn) Close() error { - if c == nil || c.Conn == nil { - return net.ErrClosed +func (c *homeDispatchConn) untrack() { + if c == nil { + return } c.once.Do(func() { if c.client != nil { @@ -593,9 +593,23 @@ func (c *homeDispatchConn) Close() error { c.client.mu.Unlock() } }) +} + +func (c *homeDispatchConn) Close() error { + if c == nil || c.Conn == nil { + return net.ErrClosed + } + c.untrack() return c.Conn.Close() } +func (c *homeDispatchConn) NetConn() net.Conn { + if c == nil { + return nil + } + return c.Conn +} + func cloneRedisOptions(options *redis.Options) *redis.Options { if options == nil { return nil @@ -1697,15 +1711,46 @@ func newPluginSyncCancelableConn(ctx context.Context, conn net.Conn) net.Conn { go func() { select { case <-ctx.Done(): - if errDeadline := conn.SetDeadline(time.Now()); errDeadline != nil { - _ = conn.Close() - } + _ = closeUnderlyingTransport(conn) case <-wrapped.done: } }() return wrapped } +func closeUnderlyingTransport(conn net.Conn) error { + if conn == nil { + return net.ErrClosed + } + current := conn + for { + if dispatchConn, ok := current.(*homeDispatchConn); ok { + dispatchConn.untrack() + if next := dispatchConn.NetConn(); next != nil && next != current { + current = next + continue + } + } + if tlsConn, ok := current.(*tls.Conn); ok { + if netConn := tlsConn.NetConn(); netConn != nil && netConn != current { + current = netConn + continue + } + } + type unwrapper interface { + NetConn() net.Conn + } + if u, ok := current.(unwrapper); ok { + if next := u.NetConn(); next != nil && next != current { + current = next + continue + } + } + break + } + return current.Close() +} + func (c *pluginSyncCancelableConn) Close() error { if c == nil || c.Conn == nil { return net.ErrClosed diff --git a/internal/home/client_test.go b/internal/home/client_test.go index effe8b43f..3c5496889 100644 --- a/internal/home/client_test.go +++ b/internal/home/client_test.go @@ -3,12 +3,17 @@ package home import ( "bufio" "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" "crypto/tls" "crypto/x509" + "crypto/x509/pkix" "encoding/json" "errors" "fmt" "io" + "math/big" "net" "net/http" "net/url" @@ -1028,6 +1033,133 @@ func TestProcessPluginSyncCommandCancellationInterruptsTLSHandshake(t *testing.T } } +func newHomeTestCertificate(t *testing.T) tls.Certificate { + t.Helper() + key, errKey := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if errKey != nil { + t.Fatalf("generate test key: %v", errKey) + } + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "127.0.0.1"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + IsCA: true, + } + der, errCreate := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + if errCreate != nil { + t.Fatalf("create test certificate: %v", errCreate) + } + leaf, errParse := x509.ParseCertificate(der) + if errParse != nil { + t.Fatalf("parse test certificate: %v", errParse) + } + return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key, Leaf: leaf} +} + +func TestProcessPluginSyncCommandCancellationUnderTLSBackpressure(t *testing.T) { + cert := newHomeTestCertificate(t) + serverTLS := &tls.Config{ + Certificates: []tls.Certificate{cert}, + MinVersion: tls.VersionTLS12, + } + rawListener, errListen := net.Listen("tcp", "127.0.0.1:0") + if errListen != nil { + t.Fatalf("listen: %v", errListen) + } + defer func() { _ = rawListener.Close() }() + + listener := tls.NewListener(rawListener, serverTLS) + defer func() { _ = listener.Close() }() + + requestReceived := make(chan struct{}) + release := make(chan struct{}) + var releaseOnce sync.Once + safeRelease := func() { + releaseOnce.Do(func() { close(release) }) + } + defer safeRelease() + + serverDone := make(chan error, 1) + + go func() { + conn, errAccept := listener.Accept() + if errAccept != nil { + serverDone <- errAccept + return + } + defer func() { _ = conn.Close() }() + reader := bufio.NewReader(conn) + _, errRead := readRedisCommand(reader) + if errRead != nil { + serverDone <- errRead + return + } + close(requestReceived) + <-release + serverDone <- nil + }() + + client := New(config.HomeConfig{Enabled: true, Host: "127.0.0.1", Port: 1, DisableClusterDiscovery: true}) + options := &redis.Options{ + Addr: rawListener.Addr().String(), + TLSConfig: &tls.Config{MinVersion: tls.VersionTLS12, InsecureSkipVerify: true}, //nolint:gosec -- test server with self-signed certificate. + DialTimeout: time.Second, + ReadTimeout: homePluginSyncOperationTimeout, + WriteTimeout: homeRedisTestOperationTimeout, + MaxRetries: -1, + ContextTimeoutEnabled: true, + } + options.Dialer = client.trackedRedisDialer(redis.NewDialer(options)) + client.cmdOptions = cloneRedisOptions(options) + client.cmd = redis.NewClient(options) + client.sub = redis.NewClient(cloneRedisOptions(options)) + t.Cleanup(func() { client.Close() }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go func() { + select { + case <-requestReceived: + cancel() + case errServer := <-serverDone: + // Push error back so post-test assertion can inspect it. + serverDone <- errServer + cancel() + } + }() + + startedAt := time.Now() + _, errSync := client.GetPluginSync(ctx, pluginstore.PluginSyncRequest{}) + safeRelease() + + select { + case errServer := <-serverDone: + if errServer != nil { + t.Fatalf("server error: %v", errServer) + } + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for server goroutine to exit") + } + + if !errors.Is(errSync, context.Canceled) { + t.Fatalf("GetPluginSync() error = %v, want context.Canceled", errSync) + } + if elapsed := time.Since(startedAt); elapsed > time.Second { + t.Fatalf("TLS backpressure cancellation took %s, want < 1s", elapsed) + } + + client.mu.Lock() + remaining := len(client.connections) + client.mu.Unlock() + if remaining != 0 { + t.Fatalf("tracked connection count = %d, want 0 after cancellation", remaining) + } +} + func TestGetPluginTasksRetainsBaseTimeout(t *testing.T) { client, _ := newRedisCommandTestClient(t, func(args []string) string { if len(args) >= 2 && args[1] == redisKeyPluginTasks {