Support TCP protocol in tshd (#14301)

This commit is contained in:
Grzegorz Zdunek
2022-07-26 14:12:56 +00:00
committed by GitHub
parent 08dcdcd27b
commit d2d9fd4dad
7 changed files with 124 additions and 17 deletions
+99 -8
View File
@@ -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
}
+6
View File
@@ -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")
}
+9 -4
View File
@@ -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 {
+1 -2
View File
@@ -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 {
+5 -3
View File
@@ -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())
+1
View File
@@ -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,
})
+3
View File
@@ -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.")