use TLS settings when setting up GRPC server or client

This commit is contained in:
Matt Jaffee 2019-10-30 17:38:48 -05:00
parent 468cf98811
commit 6f21887259
No known key found for this signature in database
GPG key ID: 08A3DFFF987B11BF
3 changed files with 34 additions and 14 deletions

View file

@ -16,10 +16,12 @@ package client
import (
"context"
"crypto/tls"
pb "github.com/pilosa/pilosa/v2/proto"
"github.com/pkg/errors"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials"
)
// GRPCClient is a client for working with the gRPC server.
@ -28,9 +30,15 @@ type GRPCClient struct {
}
// NewGRPCClient returns a new instance of GRPCClient.
func NewGRPCClient(dialTarget string) (*GRPCClient, error) {
func NewGRPCClient(dialTarget string, tlsConfig *tls.Config) (*GRPCClient, error) {
var opts []grpc.DialOption
opts = append(opts, grpc.WithInsecure()) // TODO: consider implementing WithTransportCredentials()
if tlsConfig != nil {
creds := credentials.NewTLS(tlsConfig)
opts = append(opts, grpc.WithTransportCredentials(creds))
} else {
opts = append(opts, grpc.WithInsecure())
}
gconn, err := grpc.Dial(dialTarget, opts...)
if err != nil {
return nil, errors.Wrap(err, "creating new grpc client")

View file

@ -16,6 +16,7 @@ package server
import (
"context"
"crypto/tls"
"fmt"
"log"
"net"
@ -25,6 +26,7 @@ import (
"github.com/pkg/errors"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/reflection"
"google.golang.org/grpc/status"
)
@ -556,8 +558,9 @@ func makeItems(p pilosa.RowIdentifiers) *pb.IdsOrKeys {
}
type grpcServer struct {
api *pilosa.API
hostPort string
api *pilosa.API
grpcServer *grpc.Server
hostPort string
}
type grpcServerOption func(s *grpcServer) error
@ -577,7 +580,7 @@ func OptGRPCServerURI(uri *pilosa.URI) grpcServerOption {
}
}
func (s *grpcServer) Serve() error {
func (s *grpcServer) Serve(tlsConfig *tls.Config) error {
// create listener
lis, err := net.Listen("tcp", s.hostPort)
if err != nil {
@ -585,15 +588,24 @@ func (s *grpcServer) Serve() error {
}
log.Printf("enabled grpc listening on %s", s.hostPort)
opts := make([]grpc.ServerOption, 0)
if tlsConfig != nil {
creds := credentials.NewTLS(tlsConfig)
if err != nil {
log.Fatalf("loading tls: %s\n", err)
}
opts = append(opts, grpc.Creds(creds))
}
// create grpc server
srv := grpc.NewServer()
pb.RegisterPilosaServer(srv, grpcHandler{api: s.api})
s.grpcServer = grpc.NewServer(opts...)
pb.RegisterPilosaServer(s.grpcServer, grpcHandler{api: s.api})
// register the server so its services are available to grpc_cli and others
reflection.Register(srv)
reflection.Register(s.grpcServer)
// and start...
if err := srv.Serve(lis); err != nil {
if err := s.grpcServer.Serve(lis); err != nil {
log.Fatalf("failed to serve: %v", err)
}
return nil

View file

@ -84,6 +84,7 @@ type Command struct {
API *pilosa.API
ln net.Listener
listenURI *pilosa.URI
tlsConfig *tls.Config
closeTimeout time.Duration
serverOptions []pilosa.ServerOption
@ -165,7 +166,7 @@ func (m *Command) Start() (err error) {
m.logger.Printf("listening as %s\n", m.listenURI)
go func() {
if err := m.grpcServer.Serve(); err != nil {
if err := m.grpcServer.Serve(m.tlsConfig); err != nil {
m.logger.Printf("grpc server error: %v", err)
}
}()
@ -263,9 +264,8 @@ func (m *Command) SetupServer() error {
}
// Setup TLS
var TLSConfig *tls.Config
if uri.Scheme == "https" {
TLSConfig, err = GetTLSConfig(&m.Config.TLS, m.logger.Logger())
m.tlsConfig, err = GetTLSConfig(&m.Config.TLS, m.logger.Logger())
if err != nil {
return errors.Wrap(err, "get tls config")
}
@ -281,7 +281,7 @@ func (m *Command) SetupServer() error {
return errors.Wrap(err, "new stats client")
}
m.ln, err = getListener(*uri, TLSConfig)
m.ln, err = getListener(*uri, m.tlsConfig)
if err != nil {
return errors.Wrap(err, "getting listener")
}
@ -294,7 +294,7 @@ func (m *Command) SetupServer() error {
// Save listenURI for later reference.
m.listenURI = uri
c := http.GetHTTPClient(TLSConfig)
c := http.GetHTTPClient(m.tlsConfig)
// Get advertise address as uri.
advertiseURI, err := pilosa.AddressWithDefaults(m.Config.Advertise)