diff --git a/api/client/grpc.go b/api/client/grpc.go index 6d8247d8f..5eb8d5c61 100644 --- a/api/client/grpc.go +++ b/api/client/grpc.go @@ -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") diff --git a/server/grpc.go b/server/grpc.go index 92a8febc1..8f56b95ba 100644 --- a/server/grpc.go +++ b/server/grpc.go @@ -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 diff --git a/server/server.go b/server/server.go index 64d20ce04..d1492afbd 100644 --- a/server/server.go +++ b/server/server.go @@ -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)