diff --git a/plugin/grpc/grpc.go b/plugin/grpc/grpc.go index 31c3f7de8..c6ee88ea6 100644 --- a/plugin/grpc/grpc.go +++ b/plugin/grpc/grpc.go @@ -139,6 +139,15 @@ func (g *GRPC) Name() string { return "grpc" } // Len returns the number of configured proxies. func (g *GRPC) len() int { return len(g.proxies) } +// OnShutdown closes all configured upstream connections. +func (g *GRPC) OnShutdown() error { + var err error + for _, p := range g.proxies { + err = errors.Join(err, p.close()) + } + return err +} + func (g *GRPC) match(state request.Request) bool { if !plugin.Name(g.from).Matches(state.Name()) || !g.isAllowedDomain(state.Name()) { return false diff --git a/plugin/grpc/proxy.go b/plugin/grpc/proxy.go index fc06a5a46..9b706fed0 100644 --- a/plugin/grpc/proxy.go +++ b/plugin/grpc/proxy.go @@ -37,6 +37,7 @@ type Proxy struct { addr string // connection + conn *grpc.ClientConn client pb.DnsServiceClient dialOpts []grpc.DialOption } @@ -66,11 +67,19 @@ func newProxy(addr string, tlsConfig *tls.Config) (*Proxy, error) { if err != nil { return nil, err } + p.conn = conn p.client = pb.NewDnsServiceClient(conn) return p, nil } +func (p *Proxy) close() error { + if p.conn == nil { + return nil + } + return p.conn.Close() +} + // query sends the request and waits for a response. func (p *Proxy) query(ctx context.Context, req *dns.Msg) (*dns.Msg, error) { start := time.Now() diff --git a/plugin/grpc/proxy_test.go b/plugin/grpc/proxy_test.go index 7b491634c..9425958e4 100644 --- a/plugin/grpc/proxy_test.go +++ b/plugin/grpc/proxy_test.go @@ -8,8 +8,10 @@ import ( "runtime" "slices" "testing" + "time" "github.com/coredns/caddy" + "github.com/coredns/coredns/core/dnsserver" "github.com/coredns/coredns/pb" "github.com/coredns/coredns/plugin/pkg/dnstest" "github.com/coredns/coredns/plugin/test" @@ -138,6 +140,62 @@ func TestProxyUnix(t *testing.T) { } } +func TestShutdownClosesClientConnection(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := grpc.NewServer() + pb.RegisterDnsServiceServer(server, &grpcDnsServiceServer{}) + go server.Serve(listener) + t.Cleanup(func() { + server.Stop() + listener.Close() + }) + + oldDirectives, oldCaddyQuiet, oldDNSQuiet := dnsserver.Directives, caddy.Quiet, dnsserver.Quiet + t.Cleanup(func() { + dnsserver.Directives, caddy.Quiet, dnsserver.Quiet = oldDirectives, oldCaddyQuiet, oldDNSQuiet + }) + if err := dnsserver.SetDirectives([]string{"grpc"}); err != nil { + t.Fatal(err) + } + caddy.Quiet, dnsserver.Quiet = true, true + + instance, err := caddy.Start(caddy.CaddyfileInput{ + Filepath: "Corefile", + Contents: []byte(".:0 {\ngrpc . " + listener.Addr().String() + "\n}\n"), + ServerTypeName: "dns", + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if err := instance.Stop(); err != nil { + t.Errorf("stop CoreDNS instance: %v", err) + } + instance.Wait() + }) + + query := new(dns.Msg) + query.SetQuestion("example.org.", dns.TypeA) + response, _, err := (&dns.Client{Timeout: time.Second}).Exchange(query, instance.Servers()[0].LocalAddr().String()) + if err != nil { + t.Fatalf("query before shutdown: %v", err) + } + if response.Rcode != dns.RcodeSuccess { + t.Fatalf("query before shutdown returned %s", dns.RcodeToString[response.Rcode]) + } + + if err := errors.Join(instance.ShutdownCallbacks()...); err != nil { + t.Fatalf("shutdown callbacks: %v", err) + } + response, _, err = (&dns.Client{Timeout: time.Second}).Exchange(query, instance.Servers()[0].LocalAddr().String()) + if err == nil && response.Rcode == dns.RcodeSuccess { + t.Fatal("gRPC client connection remained usable after shutdown callbacks") + } +} + type grpcDnsServiceServer struct { pb.UnimplementedDnsServiceServer } diff --git a/plugin/grpc/setup.go b/plugin/grpc/setup.go index ab72eb8ba..e1d172214 100644 --- a/plugin/grpc/setup.go +++ b/plugin/grpc/setup.go @@ -28,6 +28,7 @@ func setup(c *caddy.Controller) error { g.Next = next // Set the Next field, so the plugin chaining works. return g }) + c.OnShutdown(g.OnShutdown) return nil }