mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Support TCP protocol in tshd (#14301)
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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.")
|
||||
|
||||
Reference in New Issue
Block a user