mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-09-03 06:35:00 +08:00
feat(home): enhance homeDispatchConn with untrack method and improve connection closure handling
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user