From d2d9fd4dad9312bd94f2cf9d10fc10ae229eaa06 Mon Sep 17 00:00:00 2001 From: Grzegorz Zdunek Date: Tue, 26 Jul 2022 16:12:56 +0200 Subject: [PATCH] Support TCP protocol in tshd (#14301) --- lib/teleterm/apiserver/apiserver.go | 107 +++++++++++++++++++++++++--- lib/teleterm/apiserver/config.go | 6 ++ lib/teleterm/config.go | 13 ++-- lib/teleterm/teleterm.go | 3 +- lib/teleterm/teleterm_test.go | 8 ++- tool/tsh/daemon.go | 1 + tool/tsh/tsh.go | 3 + 7 files changed, 124 insertions(+), 17 deletions(-) diff --git a/lib/teleterm/apiserver/apiserver.go b/lib/teleterm/apiserver/apiserver.go index 414390596ca..6697a759ed9 100644 --- a/lib/teleterm/apiserver/apiserver.go +++ b/lib/teleterm/apiserver/apiserver.go @@ -15,15 +15,30 @@ package apiserver import ( + "crypto/tls" + "crypto/x509" + "fmt" "net" - "net/url" + "os" + "path/filepath" api "github.com/gravitational/teleport/lib/teleterm/api/protogen/golang/v1" "github.com/gravitational/teleport/lib/teleterm/apiserver/handler" + "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" "google.golang.org/grpc" + "google.golang.org/grpc/credentials" + + log "github.com/sirupsen/logrus" +) + +const ( + // Server certificate file name (created by tsh), Connect expects exactly the same name + tshServerCertFileName = "tsh_server.crt" + // Client certificate file name (created by Connect) + clientCertFileName = "client.crt" ) // New creates an instance of API Server @@ -46,7 +61,12 @@ func New(cfg Config) (*APIServer, error) { return nil, trace.Wrap(err) } - grpcServer := grpc.NewServer(grpc.Creds(nil), grpc.ChainUnaryInterceptor( + grpcCredentials, err := getGrpcCredentials(cfg) + if err != nil { + return nil, trace.Wrap(err) + } + + grpcServer := grpc.NewServer(grpcCredentials, grpc.ChainUnaryInterceptor( withErrorHandling(cfg.Log), )) @@ -66,24 +86,30 @@ func (s *APIServer) Stop() { } func newListener(hostAddr string) (net.Listener, error) { - uri, err := url.Parse(hostAddr) + uri, err := utils.ParseAddr(hostAddr) if err != nil { return nil, trace.BadParameter("invalid host address: %s", hostAddr) } - if uri.Scheme != "unix" { - return nil, trace.BadParameter("invalid unix socket address: %s", hostAddr) - } - - lis, err := net.Listen(uri.Scheme, uri.Path) + lis, err := net.Listen(uri.Network(), uri.Addr) if err != nil { return nil, trace.Wrap(err) } + addr := utils.FromAddr(lis.Addr()) + sendBoundNetworkPortToStdout(addr) + + log.Infof("tsh daemon is listening on %v.", addr.FullAddress()) + return lis, nil } +func sendBoundNetworkPortToStdout(addr utils.NetAddr) { + // Connect needs this message to know which port has been assigned to the server. + fmt.Printf("{CONNECT_GRPC_PORT: %v}\n", addr.Port(1)) +} + // Server is a combination of the underlying grpc.Server and its RuntimeOpts. type APIServer struct { Config @@ -92,3 +118,68 @@ type APIServer struct { // grpc is an instance of grpc server grpcServer *grpc.Server } + +func getGrpcCredentials(cfg Config) (grpc.ServerOption, error) { + uri, err := utils.ParseAddr(cfg.HostAddr) + + if err != nil { + return nil, trace.BadParameter("invalid host address: %s", cfg.HostAddr) + } + + if uri.Network() != "unix" { + keyPair, err := generateKeyPair(cfg.CertsDir) + if err != nil { + return nil, trace.Wrap(err) + } + + return grpc.Creds(keyPair), nil + } + + return grpc.Creds(nil), nil +} + +func generateKeyPair(certsDir string) (credentials.TransportCredentials, error) { + // File is first saved using under `tshServerCertTempPath` and then renamed to `tshServerCertFullPath`. + // It prevents Connect from reading half written file. + tshServerCertFullPath := filepath.Join(certsDir, tshServerCertFileName) + tshServerCertTempPath := tshServerCertFullPath + ".tmp" + + cert, err := utils.GenerateSelfSignedCert([]string{"localhost"}) + if err != nil { + return nil, trace.Wrap(err, "failed to generate a certificate") + } + + err = os.WriteFile(tshServerCertTempPath, cert.Cert, 0600) + if err != nil { + return nil, trace.Wrap(err, "failed to save server certificate") + } + + err = os.Rename(tshServerCertTempPath, tshServerCertFullPath) + if err != nil { + return nil, trace.Wrap(err, "failed to rename server certificate") + } + + certificate, err := tls.X509KeyPair(cert.Cert, cert.PrivateKey) + if err != nil { + return nil, trace.Wrap(err, "failed to parse server certificates") + } + + tlsConfig := &tls.Config{ + GetConfigForClient: func(info *tls.ClientHelloInfo) (*tls.Config, error) { + caCert, err := os.ReadFile(filepath.Join(certsDir, clientCertFileName)) + if err != nil { + return nil, trace.Wrap(err, "failed to read client certificate file") + } + caPool := x509.NewCertPool() + if !caPool.AppendCertsFromPEM(caCert) { + return nil, trace.Wrap(err, "failed to add client CA file") + } + return &tls.Config{ + ClientAuth: tls.RequireAndVerifyClientCert, + Certificates: []tls.Certificate{certificate}, + ClientCAs: caPool, + }, nil + }, + } + return credentials.NewTLS(tlsConfig), nil +} diff --git a/lib/teleterm/apiserver/config.go b/lib/teleterm/apiserver/config.go index 388090d509a..9581f5c39b6 100644 --- a/lib/teleterm/apiserver/config.go +++ b/lib/teleterm/apiserver/config.go @@ -30,6 +30,8 @@ type Config struct { Daemon *daemon.Service // Log is a component logger Log logrus.FieldLogger + // Directory containing certs used to create secure gRPC connection with daemon service + CertsDir string } // CheckAndSetDefaults checks and sets default config values. @@ -38,6 +40,10 @@ func (c *Config) CheckAndSetDefaults() error { return trace.BadParameter("missing HostAddr") } + if c.HostAddr == "" { + return trace.BadParameter("missing certs dir") + } + if c.Daemon == nil { return trace.BadParameter("missing daemon service") } diff --git a/lib/teleterm/config.go b/lib/teleterm/config.go index 8f2bf51211d..3578defd940 100644 --- a/lib/teleterm/config.go +++ b/lib/teleterm/config.go @@ -15,7 +15,6 @@ package teleterm import ( - "fmt" "os" "syscall" @@ -32,6 +31,8 @@ type Config struct { ShutdownSignals []os.Signal // HomeDir is the directory to store cluster profiles HomeDir string + // Directory containing certs used to create secure gRPC connection with daemon service + CertsDir string // InsecureSkipVerify is an option to skip HTTPS cert check InsecureSkipVerify bool } @@ -42,8 +43,12 @@ func (c *Config) CheckAndSetDefaults() error { return trace.BadParameter("missing home directory") } + if c.CertsDir == "" { + return trace.BadParameter("missing certs directory") + } + if c.Addr == "" { - c.Addr = fmt.Sprintf("unix://%v/tshd.socket", c.HomeDir) + return trace.BadParameter("missing network address") } addr, err := utils.ParseAddr(c.Addr) @@ -51,8 +56,8 @@ func (c *Config) CheckAndSetDefaults() error { return trace.Wrap(err) } - if addr.Network() != "unix" { - return trace.BadParameter("only unix sockets are supported") + if !(addr.Network() == "unix" || addr.Network() == "tcp") { + return trace.BadParameter("network address should start with unix:// or tcp:// or be empty (tcp:// is used in that case)") } if len(c.ShutdownSignals) == 0 { diff --git a/lib/teleterm/teleterm.go b/lib/teleterm/teleterm.go index ddcf47992f2..906eda902b9 100644 --- a/lib/teleterm/teleterm.go +++ b/lib/teleterm/teleterm.go @@ -52,6 +52,7 @@ func Serve(ctx context.Context, cfg Config) error { apiServer, err := apiserver.New(apiserver.Config{ HostAddr: cfg.Addr, Daemon: daemonService, + CertsDir: cfg.CertsDir, }) if err != nil { return trace.Wrap(err) @@ -77,8 +78,6 @@ func Serve(ctx context.Context, cfg Config) error { apiServer.Stop() }() - log.Infof("tsh daemon is listening on %v.", cfg.Addr) - errAPI := <-serverAPIWait if errAPI != nil { diff --git a/lib/teleterm/teleterm_test.go b/lib/teleterm/teleterm_test.go index 778e00c1151..3a26487927b 100644 --- a/lib/teleterm/teleterm_test.go +++ b/lib/teleterm/teleterm_test.go @@ -25,16 +25,18 @@ import ( "time" "github.com/gravitational/teleport/lib/teleterm" - "github.com/stretchr/testify/require" ) func TestStart(t *testing.T) { homeDir := t.TempDir() + certsDir := t.TempDir() sockPath := filepath.Join(homeDir, "teleterm.sock") + cfg := teleterm.Config{ - Addr: fmt.Sprintf("unix://%v", sockPath), - HomeDir: fmt.Sprintf("%v/", homeDir), + Addr: fmt.Sprintf("unix://%v", sockPath), + HomeDir: homeDir, + CertsDir: certsDir, } ctx, cancel := context.WithCancel(context.Background()) diff --git a/tool/tsh/daemon.go b/tool/tsh/daemon.go index 5b63b825f5e..7c38275a373 100644 --- a/tool/tsh/daemon.go +++ b/tool/tsh/daemon.go @@ -41,6 +41,7 @@ func onDaemonStart(cf *CLIConf) error { err := teleterm.Serve(ctx, teleterm.Config{ HomeDir: homeDir, + CertsDir: cf.DaemonCertsDir, Addr: cf.DaemonAddr, InsecureSkipVerify: cf.InsecureSkipVerify, }) diff --git a/tool/tsh/tsh.go b/tool/tsh/tsh.go index 5b311e7c18f..c9645f827fb 100644 --- a/tool/tsh/tsh.go +++ b/tool/tsh/tsh.go @@ -176,6 +176,8 @@ type CLIConf struct { KubernetesCluster string // DaemonAddr is the daemon listening address. DaemonAddr string + // DaemonCertsDir is the directory containing certs used to create secure gRPC connection with daemon service + DaemonCertsDir string // DatabaseService specifies the database proxy server to log into. DatabaseService string // DatabaseUser specifies database user to embed in the certificate. @@ -532,6 +534,7 @@ func Run(ctx context.Context, args []string, opts ...cliOption) error { daemon := app.Command("daemon", "Daemon is the tsh daemon service").Hidden() daemonStart := daemon.Command("start", "Starts tsh daemon service").Hidden() daemonStart.Flag("addr", "Addr is the daemon listening address.").StringVar(&cf.DaemonAddr) + daemonStart.Flag("certs-dir", "Directory containing certs used to create secure gRPC connection with daemon service").StringVar(&cf.DaemonCertsDir) // AWS. aws := app.Command("aws", "Access AWS API.")