diff --git a/api/client/contextdialer.go b/api/client/contextdialer.go index 3106e932018..da7a3ff9308 100644 --- a/api/client/contextdialer.go +++ b/api/client/contextdialer.go @@ -121,7 +121,7 @@ func newTLSRoutingTunnelDialer(ssh ssh.ClientConfig, keepAlivePeriod, dialTimeou } - host, err := webclient.ExtractHost(tunnelAddr) + host, _, err := webclient.ParseHostPort(tunnelAddr) if err != nil { return nil, trace.Wrap(err) } diff --git a/api/client/webclient/webclient.go b/api/client/webclient/webclient.go index bb33154001e..411e461fba5 100644 --- a/api/client/webclient/webclient.go +++ b/api/client/webclient/webclient.go @@ -222,7 +222,7 @@ func GetTunnelAddr(cfg *Config) (string, error) { } // If TELEPORT_TUNNEL_PUBLIC_ADDR is set, nothing else has to be done, return it. if tunnelAddr := os.Getenv(defaults.TunnelPublicAddrEnvar); tunnelAddr != "" { - return extractHostPort(tunnelAddr) + return parseAndJoinHostPort(tunnelAddr) } // Ping web proxy to retrieve tunnel proxy address. @@ -230,7 +230,13 @@ func GetTunnelAddr(cfg *Config) (string, error) { if err != nil { return "", trace.Wrap(err) } - return tunnelAddr(cfg.ProxyAddr, pr.Proxy) + // DELETE IN 11.0.0 + // newer proxies should return WebListenAddr so + // we don't need to rely on the dialed proxyAddr + if pr.Proxy.SSH.WebListenAddr == "" { + pr.Proxy.SSH.WebListenAddr = cfg.ProxyAddr + } + return pr.Proxy.tunnelProxyAddr() } func GetMOTD(cfg *Config) (*MotD, error) { @@ -281,6 +287,8 @@ type PingResponse struct { ServerVersion string `json:"server_version"` // MinClientVersion is the minimum client version required by the server. MinClientVersion string `json:"min_client_version"` + // ClusterName contains the name of the Teleport cluster. + ClusterName string `json:"cluster_name"` } // PingErrorResponse contains the error message if the requested connector @@ -328,6 +336,9 @@ type SSHProxySettings struct { // listening for connections on. TunnelListenAddr string `json:"tunnel_listen_addr,omitempty"` + // WebListenAddr is the address where the proxy web handler is listening. + WebListenAddr string `json:"web_listen_addr,omitempty"` + // PublicAddr is the public address of the HTTP proxy. PublicAddr string `json:"public_addr,omitempty"` @@ -428,131 +439,156 @@ type GithubSettings struct { Display string `json:"display"` } -// The tunnel addr is retrieved in the following preference order: -// 1. If proxy support ALPN listener where all services are exposed on single port return ProxyPublicAddr/ProxyAddr. -// 2. Reverse Tunnel Public Address. -// 3. SSH Proxy Public Address Host + Tunnel Port. -// 4. HTTP Proxy Public Address Host + Tunnel Port. -// 5. Proxy Address Host + Tunnel Port. -func tunnelAddr(proxyAddr string, settings ProxySettings) (string, error) { - if settings.TLSRoutingEnabled { - return tunnelAddrForTLSRouting(proxyAddr, settings) - } - - // If a tunnel public address is set, nothing else has to be done, return it. - sshSettings := settings.SSH - if sshSettings.TunnelPublicAddr != "" { - return extractHostPort(sshSettings.TunnelPublicAddr) - } - - // Extract the port the tunnel server is listening on. - tunnelPort := strconv.Itoa(defaults.SSHProxyTunnelListenPort) - if sshSettings.TunnelListenAddr != "" { - if port, err := extractPort(sshSettings.TunnelListenAddr); err == nil { - tunnelPort = port +// tunnelProxyAddr returns the tunnel proxy address for the proxy settings. +func (ps *ProxySettings) tunnelProxyAddr() (string, error) { + if ps.TLSRoutingEnabled { + webPort := ps.getWebPort() + switch { + case ps.SSH.PublicAddr != "": + return parseAndJoinHostPort(ps.SSH.PublicAddr, WithDefaultPort(webPort)) + default: + return parseAndJoinHostPort(ps.SSH.WebListenAddr, WithDefaultPort(webPort)) } } - // If a tunnel public address has not been set, but a related HTTP or SSH - // public address has been set, extract the hostname but use the port from - // the tunnel listen address. - if sshSettings.SSHPublicAddr != "" { - if host, err := ExtractHost(sshSettings.SSHPublicAddr); err == nil { - return net.JoinHostPort(host, tunnelPort), nil - } + tunnelPort := ps.getTunnelPort() + switch { + case ps.SSH.TunnelPublicAddr != "": + return parseAndJoinHostPort(ps.SSH.TunnelPublicAddr, WithDefaultPort(tunnelPort)) + case ps.SSH.SSHPublicAddr != "": + return parseAndJoinHostPort(ps.SSH.SSHPublicAddr, WithOverridePort(tunnelPort)) + case ps.SSH.PublicAddr != "": + return parseAndJoinHostPort(ps.SSH.PublicAddr, WithOverridePort(tunnelPort)) + case ps.SSH.TunnelListenAddr != "": + return parseAndJoinHostPort(ps.SSH.TunnelListenAddr, WithDefaultPort(tunnelPort)) + default: + // If nothing else is set, we can at least try the WebListenAddr which should always be set + return parseAndJoinHostPort(ps.SSH.WebListenAddr, WithDefaultPort(tunnelPort)) } - if sshSettings.PublicAddr != "" { - if host, err := ExtractHost(sshSettings.PublicAddr); err == nil { - return net.JoinHostPort(host, tunnelPort), nil - } - } - - // If nothing is set, fallback to the address dialed with tunnel port. - host, err := ExtractHost(proxyAddr) - if err != nil { - return "", trace.Wrap(err, "failed to parse the given proxy address") - } - return net.JoinHostPort(host, tunnelPort), nil } -// tunnelAddrForTLSRouting returns reverse tunnel proxy address for proxy supporting TLS Routing. -func tunnelAddrForTLSRouting(proxyAddr string, settings ProxySettings) (string, error) { - if settings.SSH.PublicAddr != "" { - // Check if PublicAddr contains a port number. - if _, err := extractPort(settings.SSH.PublicAddr); err == nil { - return extractHostPort(settings.SSH.PublicAddr) - } - // Get port number from proxyAddr or use default one. - port := strconv.Itoa(defaults.ProxyWebListenPort) - if webPort, err := extractPort(proxyAddr); err == nil { - port = webPort - } - - if host, err := ExtractHost(settings.SSH.PublicAddr); err == nil { - return net.JoinHostPort(host, port), nil +// SSHProxyHostPort returns the ssh proxy host and port for the proxy settings. +func (ps *ProxySettings) SSHProxyHostPort() (host, port string, err error) { + if ps.TLSRoutingEnabled { + webPort := ps.getWebPort() + switch { + case ps.SSH.PublicAddr != "": + return ParseHostPort(ps.SSH.PublicAddr, WithDefaultPort(webPort)) + default: + return ParseHostPort(ps.SSH.WebListenAddr, WithDefaultPort(webPort)) } } - // Got proxyAddr with a port number for instance: proxy.example.com:3080 - if _, err := extractPort(proxyAddr); err == nil { - return proxyAddr, nil + sshPort := ps.getSSHPort() + switch { + case ps.SSH.SSHPublicAddr != "": + return ParseHostPort(ps.SSH.SSHPublicAddr, WithDefaultPort(sshPort)) + case ps.SSH.PublicAddr != "": + return ParseHostPort(ps.SSH.PublicAddr, WithOverridePort(sshPort)) + case ps.SSH.ListenAddr != "": + return ParseHostPort(ps.SSH.ListenAddr, WithDefaultPort(sshPort)) + default: + // If nothing else is set, we can at least try the WebListenAddr which should always be set + return ParseHostPort(ps.SSH.WebListenAddr, WithDefaultPort(sshPort)) } - host, err := ExtractHost(proxyAddr) - if err != nil { - return "", trace.Wrap(err, "failed to parse the given proxy address") - } - - // Got proxy address without a port like: proxy.example.com - // If proxyAddr doesn't contain any port it means that HTTPS port should be used because during Find call - // The destination URL is constructed by the fmt.Sprintf("https://%s/webapi/find", proxyAddr) function. - return net.JoinHostPort(host, strconv.Itoa(defaults.StandardHTTPSPort)), nil } -// extractHostPort takes addresses like "tcp://host:port/path" and returns "host:port". -func extractHostPort(addr string) (string, error) { +// getWebPort from WebListenAddr or global default +func (ps *ProxySettings) getWebPort() int { + if webPort, err := parsePort(ps.SSH.WebListenAddr); err == nil { + return webPort + } + return defaults.StandardHTTPSPort +} + +// getSSHPort from ListenAddr or global default +func (ps *ProxySettings) getSSHPort() int { + if webPort, err := parsePort(ps.SSH.ListenAddr); err == nil { + return webPort + } + return defaults.SSHProxyListenPort +} + +// getTunnelPort from TunnelListenAddr or global default +func (ps *ProxySettings) getTunnelPort() int { + if webPort, err := parsePort(ps.SSH.TunnelListenAddr); err == nil { + return webPort + } + return defaults.SSHProxyTunnelListenPort +} + +type ParseHostPortOpt func(host, port string) (hostR, portR string) + +// WithDefaultPort replaces the parse port with the default port if empty. +func WithDefaultPort(defaultPort int) ParseHostPortOpt { + defaultPortString := strconv.Itoa(defaultPort) + return func(host, port string) (string, string) { + if port == "" { + return host, defaultPortString + } + return host, port + } +} + +// WithOverridePort replaces the parsed port with the override port. +func WithOverridePort(overridePort int) ParseHostPortOpt { + overridePortString := strconv.Itoa(overridePort) + return func(host, port string) (string, string) { + return host, overridePortString + } +} + +// ParseHostPort parses host and port from the given address. +func ParseHostPort(addr string, opts ...ParseHostPortOpt) (host, port string, err error) { if addr == "" { - return "", trace.BadParameter("missing parameter address") + return "", "", trace.BadParameter("missing parameter address") } if !strings.Contains(addr, "://") { addr = "tcp://" + addr } u, err := url.Parse(addr) if err != nil { - return "", trace.BadParameter("failed to parse %q: %v", addr, err) + return "", "", trace.BadParameter("failed to parse %q: %v", addr, err) } switch u.Scheme { case "tcp", "http", "https": - return u.Host, nil default: - return "", trace.BadParameter("'%v': unsupported scheme: '%v'", addr, u.Scheme) + return "", "", trace.BadParameter("'%v': unsupported scheme: '%v'", addr, u.Scheme) } + host, port, err = net.SplitHostPort(u.Host) + if err != nil && strings.Contains(err.Error(), "missing port in address") { + host = u.Host + } else if err != nil { + return "", "", trace.Wrap(err) + } + for _, opt := range opts { + host, port = opt(host, port) + } + return host, port, nil } -// ExtractHost takes addresses like "tcp://host:port/path" and returns "host". -func ExtractHost(addr string) (ra string, err error) { - parsed, err := extractHostPort(addr) +// parseAndJoinHostPort parses host and port from the given address and returns "host:port". +func parseAndJoinHostPort(addr string, opts ...ParseHostPortOpt) (string, error) { + host, port, err := ParseHostPort(addr, opts...) if err != nil { return "", trace.Wrap(err) + } else if port == "" { + return host, nil } - host, _, err := net.SplitHostPort(parsed) - if err != nil { - if strings.Contains(err.Error(), "missing port in address") { - return addr, nil - } - return "", trace.Wrap(err) - } - return host, nil + return net.JoinHostPort(host, port), nil } -// extractPort takes addresses like "tcp://host:port/path" and returns "port". -func extractPort(addr string) (string, error) { - parsed, err := extractHostPort(addr) +// parsePort parses port from the given address as an integer. +func parsePort(addr string) (int, error) { + _, port, err := ParseHostPort(addr) if err != nil { - return "", trace.Wrap(err) + return 0, trace.Wrap(err) + } else if port == "" { + return 0, trace.BadParameter("missing port in address %q", addr) } - _, port, err := net.SplitHostPort(parsed) + portI, err := strconv.Atoi(port) if err != nil { - return "", trace.Wrap(err) + return 0, trace.Wrap(err) } - return port, nil + return portI, nil } diff --git a/api/client/webclient/webclient_test.go b/api/client/webclient/webclient_test.go index 6cbb554394d..9b0fd15abe4 100644 --- a/api/client/webclient/webclient_test.go +++ b/api/client/webclient/webclient_test.go @@ -112,7 +112,6 @@ func TestGetTunnelAddr(t *testing.T) { func TestTunnelAddr(t *testing.T) { type testCase struct { - proxyAddr string settings ProxySettings expectedTunnelAddr string } @@ -120,100 +119,103 @@ func TestTunnelAddr(t *testing.T) { testTunnelAddr := func(tc testCase) func(*testing.T) { return func(t *testing.T) { t.Parallel() - tunnelAddr, err := tunnelAddr(tc.proxyAddr, tc.settings) + tunnelAddr, err := tc.settings.tunnelProxyAddr() require.NoError(t, err) require.Equal(t, tc.expectedTunnelAddr, tunnelAddr) } } t.Run("should use TunnelPublicAddr", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com", settings: ProxySettings{ SSH: SSHProxySettings{ TunnelPublicAddr: "tunnel.example.com:4024", PublicAddr: "public.example.com", SSHPublicAddr: "ssh.example.com", TunnelListenAddr: "[::]:5024", + WebListenAddr: "proxy.example.com", }, }, expectedTunnelAddr: "tunnel.example.com:4024", })) t.Run("should use SSHPublicAddr and TunnelListenAddr", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com", settings: ProxySettings{ SSH: SSHProxySettings{ SSHPublicAddr: "ssh.example.com", PublicAddr: "public.example.com", TunnelListenAddr: "[::]:5024", + WebListenAddr: "proxy.example.com", }, }, expectedTunnelAddr: "ssh.example.com:5024", })) t.Run("should use PublicAddr and TunnelListenAddr", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com", settings: ProxySettings{ SSH: SSHProxySettings{ PublicAddr: "public.example.com", TunnelListenAddr: "[::]:5024", + WebListenAddr: "proxy.example.com", }, }, expectedTunnelAddr: "public.example.com:5024", })) t.Run("should use PublicAddr and SSHProxyTunnelListenPort", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com", settings: ProxySettings{ SSH: SSHProxySettings{ - PublicAddr: "public.example.com", + PublicAddr: "public.example.com", + WebListenAddr: "proxy.example.com", }, }, expectedTunnelAddr: "public.example.com:3024", })) - t.Run("should use proxyAddr and SSHProxyTunnelListenPort", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com", - settings: ProxySettings{SSH: SSHProxySettings{}}, + t.Run("should use WebListenAddr and SSHProxyTunnelListenPort", testTunnelAddr(testCase{ + settings: ProxySettings{ + SSH: SSHProxySettings{ + WebListenAddr: "proxy.example.com", + }, + }, expectedTunnelAddr: "proxy.example.com:3024", })) t.Run("should use PublicAddr with ProxyWebPort if TLSRoutingEnabled was enabled", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com:443", settings: ProxySettings{ SSH: SSHProxySettings{ PublicAddr: "public.example.com", TunnelListenAddr: "[::]:5024", TunnelPublicAddr: "tpa.example.com:3032", + WebListenAddr: "proxy.example.com:443", }, TLSRoutingEnabled: true, }, expectedTunnelAddr: "public.example.com:443", })) t.Run("should use PublicAddr with custom port if TLSRoutingEnabled was enabled", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com:443", settings: ProxySettings{ SSH: SSHProxySettings{ PublicAddr: "public.example.com:443", TunnelListenAddr: "[::]:5024", TunnelPublicAddr: "tpa.example.com:3032", + WebListenAddr: "proxy.example.com:443", }, TLSRoutingEnabled: true, }, expectedTunnelAddr: "public.example.com:443", })) - t.Run("should use proxyAddr with custom ProxyWebPort if TLSRoutingEnabled was enabled", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com:443", + t.Run("should use WebListenAddr with custom ProxyWebPort if TLSRoutingEnabled was enabled", testTunnelAddr(testCase{ settings: ProxySettings{ SSH: SSHProxySettings{ TunnelListenAddr: "[::]:5024", TunnelPublicAddr: "tpa.example.com:3032", + WebListenAddr: "proxy.example.com:443", }, TLSRoutingEnabled: true, }, expectedTunnelAddr: "proxy.example.com:443", })) - t.Run("should use proxyAddr with default https port if TLSRoutingEnabled was enabled", testTunnelAddr(testCase{ - proxyAddr: "proxy.example.com", + t.Run("should use WebListenAddr with default https port if TLSRoutingEnabled was enabled", testTunnelAddr(testCase{ settings: ProxySettings{ SSH: SSHProxySettings{ TunnelListenAddr: "[::]:5024", TunnelPublicAddr: "tpa.example.com:3032", + WebListenAddr: "proxy.example.com", }, TLSRoutingEnabled: true, }, @@ -221,72 +223,81 @@ func TestTunnelAddr(t *testing.T) { })) } -func TestExtract(t *testing.T) { +func TestParse(t *testing.T) { testCases := []struct { addr string hostPort string host string - port string + port int }{ { addr: "example.com", hostPort: "example.com", host: "example.com", - port: "", + port: 0, }, { addr: "example.com:443", hostPort: "example.com:443", host: "example.com", - port: "443", + port: 443, }, { addr: "http://example.com:443", hostPort: "example.com:443", host: "example.com", - port: "443", + port: 443, }, { addr: "https://example.com:443", hostPort: "example.com:443", host: "example.com", - port: "443", + port: 443, }, { addr: "tcp://example.com:443", hostPort: "example.com:443", host: "example.com", - port: "443", + port: 443, }, { addr: "file://host/path", hostPort: "", host: "", - port: "", + port: 0, }, { addr: "[::]:443", hostPort: "[::]:443", host: "::", - port: "443", + port: 443, }, { addr: "https://example.com:443/path?query=query#fragment", hostPort: "example.com:443", host: "example.com", - port: "443", + port: 443, }, } for _, tc := range testCases { t.Run(tc.addr, func(t *testing.T) { - hostPort, err := extractHostPort(tc.addr) - // Expect err if expected value is empty - require.True(t, (tc.hostPort == "") == (err != nil)) - require.Equal(t, tc.hostPort, hostPort) + hostPort, err := parseAndJoinHostPort(tc.addr) + if tc.hostPort == "" { + require.Error(t, err) + } else { + require.NoError(t, err) + require.Equal(t, tc.hostPort, hostPort) + } - host, err := ExtractHost(tc.addr) - // Expect err if expected value is empty - require.True(t, (tc.host == "") == (err != nil)) - require.Equal(t, tc.host, host) + host, _, err := ParseHostPort(tc.addr) + if tc.host == "" { + require.Error(t, err) + } else { + require.NoError(t, err) + require.Equal(t, tc.host, host) + } - port, err := extractPort(tc.addr) - // Expect err if expected value is empty - require.True(t, (tc.port == "") == (err != nil)) - require.Equal(t, tc.port, port) + port, err := parsePort(tc.addr) + if tc.port == 0 { + require.Error(t, err) + } else { + require.NoError(t, err) + require.Equal(t, tc.port, port) + } }) } } @@ -339,3 +350,92 @@ func TestNewWebClientIgnoreProxy(t *testing.T) { require.Contains(t, err.Error(), "lookup fakedomain.example.com") require.Contains(t, err.Error(), "no such host") } + +func TestSSHProxyHostPort(t *testing.T) { + tests := []struct { + testName string + inProxySettings ProxySettings + outHost string + outPort string + }{ + { + testName: "TLS routing enabled, web public addr", + inProxySettings: ProxySettings{ + SSH: SSHProxySettings{ + PublicAddr: "proxy.example.com:443", + WebListenAddr: "127.0.0.1:3080", + }, + TLSRoutingEnabled: true, + }, + outHost: "proxy.example.com", + outPort: "443", + }, + { + testName: "TLS routing enabled, web public addr with listen addr", + inProxySettings: ProxySettings{ + SSH: SSHProxySettings{ + PublicAddr: "proxy.example.com", + WebListenAddr: "127.0.0.1:443", + }, + TLSRoutingEnabled: true, + }, + outHost: "proxy.example.com", + outPort: "443", + }, + { + testName: "TLS routing enabled, web listen addr", + inProxySettings: ProxySettings{ + SSH: SSHProxySettings{ + WebListenAddr: "127.0.0.1:3080", + }, + TLSRoutingEnabled: true, + }, + outHost: "127.0.0.1", + outPort: "3080", + }, + { + testName: "TLS routing disabled, SSH public addr", + inProxySettings: ProxySettings{ + SSH: SSHProxySettings{ + SSHPublicAddr: "ssh.example.com:3023", + PublicAddr: "proxy.example.com:443", + ListenAddr: "127.0.0.1:3023", + }, + TLSRoutingEnabled: false, + }, + outHost: "ssh.example.com", + outPort: "3023", + }, + { + testName: "TLS routing disabled, web public addr", + inProxySettings: ProxySettings{ + SSH: SSHProxySettings{ + PublicAddr: "proxy.example.com:443", + ListenAddr: "127.0.0.1:3023", + }, + TLSRoutingEnabled: false, + }, + outHost: "proxy.example.com", + outPort: "3023", + }, + { + testName: "TLS routing disabled, SSH listen addr", + inProxySettings: ProxySettings{ + SSH: SSHProxySettings{ + ListenAddr: "127.0.0.1:3023", + }, + TLSRoutingEnabled: false, + }, + outHost: "127.0.0.1", + outPort: "3023", + }, + } + for _, test := range tests { + t.Run(test.testName, func(t *testing.T) { + host, port, err := test.inProxySettings.SSHProxyHostPort() + require.NoError(t, err) + require.Equal(t, test.outHost, host) + require.Equal(t, test.outPort, port) + }) + } +} diff --git a/api/defaults/defaults.go b/api/defaults/defaults.go index b38fe58707c..6441dab5aac 100644 --- a/api/defaults/defaults.go +++ b/api/defaults/defaults.go @@ -127,6 +127,9 @@ const ( // run behind an environment/firewall which only allows outgoing connections) SSHProxyTunnelListenPort = 3024 + // SSHProxyListenPort is the default Teleport SSH proxy listen port. + SSHProxyListenPort = 3023 + // ProxyWebListenPort is the default Teleport Proxy WebPort address. ProxyWebListenPort = 3080 diff --git a/lib/client/api.go b/lib/client/api.go index d27c6115837..79accbf23c8 100644 --- a/lib/client/api.go +++ b/lib/client/api.go @@ -3162,6 +3162,29 @@ func (tc *TeleportClient) ShowMOTD(ctx context.Context) error { return nil } +// UpdateKnownHosts updates ~/.tsh/known_hosts with trusted host certificate +// authorities for the specified proxy and cluster. +func (tc *TeleportClient) UpdateKnownHosts(ctx context.Context, proxyHost, clusterName string) error { + trustedCAs, err := tc.GetTrustedCA(ctx, clusterName) + if err != nil { + return trace.Wrap(err) + } + for _, ca := range auth.AuthoritiesToTrustedCerts(trustedCAs) { + if ca.ClusterName != clusterName { + continue + } + hostCerts, err := ca.SSHCertPublicKeys() + if err != nil { + return trace.Wrap(err) + } + err = tc.localAgent.keyStore.AddKnownHostKeys(clusterName, proxyHost, hostCerts) + if err != nil { + return trace.Wrap(err) + } + } + return nil +} + // GetTrustedCA returns a list of host certificate authorities // trusted by the cluster client is authenticated with. func (tc *TeleportClient) GetTrustedCA(ctx context.Context, clusterName string) ([]types.CertAuthority, error) { diff --git a/lib/client/keyagent.go b/lib/client/keyagent.go index 3b48c088509..08911522d56 100644 --- a/lib/client/keyagent.go +++ b/lib/client/keyagent.go @@ -180,6 +180,11 @@ func (a *LocalKeyAgent) UpdateUsername(username string) { a.username = username } +// UpdateCluster changes the cluster that the local agent operates on. +func (a *LocalKeyAgent) UpdateCluster(cluster string) { + a.siteName = cluster +} + // LoadKeyForCluster fetches a cluster-specific SSH key and loads it into the // SSH agent. func (a *LocalKeyAgent) LoadKeyForCluster(clusterName string) (*agent.AddedKey, error) { diff --git a/lib/service/proxy_settings.go b/lib/service/proxy_settings.go index b5671436c73..1e85d4e5a9b 100644 --- a/lib/service/proxy_settings.go +++ b/lib/service/proxy_settings.go @@ -68,6 +68,7 @@ func (p *proxySettings) buildProxySettings(proxyListenerMode types.ProxyListener SSH: webclient.SSHProxySettings{ ListenAddr: p.proxySSHAddr.String(), TunnelListenAddr: p.cfg.Proxy.ReverseTunnelListenAddr.String(), + WebListenAddr: p.cfg.Proxy.WebAddr.String(), }, } @@ -98,6 +99,7 @@ func (p *proxySettings) buildProxySettingsV2(proxyListenerMode types.ProxyListen if proxyListenerMode == types.ProxyListenerMode_Multiplex { settings.SSH.ListenAddr = multiplexAddr settings.SSH.TunnelListenAddr = multiplexAddr + settings.SSH.WebListenAddr = multiplexAddr settings.Kube.ListenAddr = multiplexAddr settings.DB.MySQLListenAddr = multiplexAddr settings.DB.PostgresListenAddr = multiplexAddr diff --git a/lib/web/apiserver.go b/lib/web/apiserver.go index 7d7f6575036..088e68b74db 100644 --- a/lib/web/apiserver.go +++ b/lib/web/apiserver.go @@ -792,6 +792,7 @@ func (h *Handler) ping(w http.ResponseWriter, r *http.Request, p httprouter.Para Proxy: *proxyConfig, ServerVersion: teleport.Version, MinClientVersion: teleport.MinClientVersion, + ClusterName: h.auth.clusterName, }, nil } @@ -804,6 +805,7 @@ func (h *Handler) find(w http.ResponseWriter, r *http.Request, p httprouter.Para Proxy: *proxyConfig, ServerVersion: teleport.Version, MinClientVersion: teleport.MinClientVersion, + ClusterName: h.auth.clusterName, }, nil } @@ -823,6 +825,7 @@ func (h *Handler) pingWithConnector(w http.ResponseWriter, r *http.Request, p ht response := &webclient.PingResponse{ Proxy: *proxyConfig, ServerVersion: teleport.Version, + ClusterName: h.auth.clusterName, } hasMessageOfTheDay := cap.GetMessageOfTheDay() != "" diff --git a/rfd/0062-tsh-proxy-template.md b/rfd/0062-tsh-proxy-template.md index 5711ac61835..e9134f9db91 100644 --- a/rfd/0062-tsh-proxy-template.md +++ b/rfd/0062-tsh-proxy-template.md @@ -1,6 +1,6 @@ --- authors: Roman Tkachenko (roman@goteleport.com) -state: draft +state: implemented --- # RFD 62 - Proxy template support for tsh proxy @@ -69,41 +69,45 @@ Specifically, the syntax for the `` would look like: Host *.acme.com HostName %h Port 3022 - ProxyCommand tsh proxy ssh --proxy={{proxy}} %r@%h:%p + ProxyCommand tsh proxy ssh -J {{proxy}} %r@%h:%p ``` -With the `--proxy` flag set, the command connects directly to the specified -proxy instead of the default behavior of connecting to the proxy of the current -client profile. +With the `-J` flag set, the command connects directly to the specified proxy +instead of the default behavior of connecting to the proxy of the current +client profile. This usage of the `-J` flag is consistent with the existing +proxy jump functionality (`tsh ssh -J`) and [Cluster Routing](https://github.com/gravitational/teleport/blob/master/rfd/0021-cluster-routing.md). -When a templating variable `{{proxy}}` is used, the host name and proxy address -are extracted from the full hostname in the `%r@%h:%p` spec. - -Users define the rules of how to parse node/proxy from the full hostname in -the tsh config file `$TELEPORT_HOME/config/config.yaml`. Similar to role -templating, group captures are supported: +When a template variable `{{proxy}}` is used, the host name and proxy address +are extracted from the full hostname in the `%r@%h:%p` spec. Users define the +rules of how to parse node/proxy from the full hostname in the tsh config file +`$TELEPORT_HOME/config/config.yaml` (or global `/etc/tsh.yaml`). Group captures +are supported: ```yaml proxy_templates: # Example template where nodes have short names like node-1, node-2, etc. -- template: "^(\w+)\.(leaf1.us.acme.com)$" +- template: '^(\w+)\.(leaf1.us.acme.com)$' host: "$1" # host is optional and will default to the full %h if not specified - proxy: "$2:3023" + proxy: "$2:3080" # Example template where nodes have FQDN names like node-1.leaf2.eu.acme.com. -- template: "^(\w+)\.(leaf2.eu.acme.com)$" +- template: '^(\w+)\.(leaf2.eu.acme.com)$' proxy: "$2:443" ``` Templates are evaluated in order and the first one matching will take effect. +Note that the proxy address must point to the web proxy address (not SSH proxy): +`tsh proxy ssh` will issue a ping request to the proxy to retrieve additional +information about the cluster, including the SSH proxy endpoint. + In the example described above, where the user has nodes `node-1`, `node-2` in multiple leaf clusters, their template configuration can look like: ```yaml proxy_templates: -- template: "^([^\.]+)\.(.+)$" +- template: '^([^\.]+)\.(.+)$' host: "$1" - proxy: "$2:3023" + proxy: "$2:3080" ``` In the node spec `%r@%h:%p` the host name `%h` will be replaced by the host from @@ -113,13 +117,13 @@ the template. So given the above proxy template configuration, the following proxy command: ```bash -tsh proxy ssh --proxy={{proxy}} %r@%h:%p +tsh proxy ssh -J {{proxy}} %r@%h:%p ``` is equivalent to the following when connecting to `node-1.leaf1.us.acme.com`: ```bash -tsh proxy ssh --proxy=leaf1.us.acme.com:3023 %r@node-1:3022 +tsh proxy ssh -J leaf1.us.acme.com:3080 %r@node-1:3022 ``` ### Auto-login diff --git a/tool/tsh/proxy.go b/tool/tsh/proxy.go index 885204e87be..f676fb61d21 100644 --- a/tool/tsh/proxy.go +++ b/tool/tsh/proxy.go @@ -31,9 +31,9 @@ import ( "github.com/gravitational/trace" + "github.com/gravitational/teleport/api/client/webclient" "github.com/gravitational/teleport/api/profile" "github.com/gravitational/teleport/api/utils/keypaths" - "github.com/gravitational/teleport/lib/client" libclient "github.com/gravitational/teleport/lib/client" "github.com/gravitational/teleport/lib/client/db/dbcmd" "github.com/gravitational/teleport/lib/defaults" @@ -55,17 +55,110 @@ func onProxyCommandSSH(cf *CLIConf) error { return trace.Wrap(err) } - targetHost, targetPort, err := net.SplitHostPort(tc.Host) + proxyParams, err := getSSHProxyParams(cf, tc) if err != nil { return trace.Wrap(err) } - targetHost = cleanTargetHost(targetHost, tc.WebProxyHost(), tc.SiteName) - if tc.TLSRoutingEnabled { - return trace.Wrap(sshProxyWithTLSRouting(cf, tc, targetHost, targetPort)) + if len(tc.JumpHosts) > 0 { + err := setupJumpHost(cf, tc, *proxyParams) + if err != nil { + return trace.Wrap(err) + } } - return trace.Wrap(sshProxy(tc, targetHost, targetPort)) + if proxyParams.tlsRouting { + return trace.Wrap(sshProxyWithTLSRouting(cf, tc, *proxyParams)) + } + + return trace.Wrap(sshProxy(tc, *proxyParams)) +} + +// getSSHProxyParams prepares parameters for establishing an SSH proxy. +func getSSHProxyParams(cf *CLIConf, tc *libclient.TeleportClient) (*sshProxyParams, error) { + targetHost, targetPort, err := net.SplitHostPort(tc.Host) + if err != nil { + return nil, trace.Wrap(err) + } + // Without jump hosts, we will be connecting to the current Teleport client + // proxy the user is logged into. + if len(tc.JumpHosts) == 0 { + proxyHost, proxyPort := tc.SSHProxyHostPort() + if tc.TLSRoutingEnabled { + proxyHost, proxyPort = tc.WebProxyHostPort() + } + return &sshProxyParams{ + proxyHost: proxyHost, + proxyPort: strconv.Itoa(proxyPort), + targetHost: cleanTargetHost(targetHost, tc.WebProxyHost(), tc.SiteName), + targetPort: targetPort, + clusterName: tc.SiteName, + tlsRouting: tc.TLSRoutingEnabled, + }, nil + } + // When jump host is specified, we will be connecting to the jump host's + // proxy directly. Call its ping endpoint to figure out the cluster details + // such as cluster name, SSH proxy address, etc. + ping, err := webclient.Find(&webclient.Config{ + Context: cf.Context, + ProxyAddr: tc.JumpHosts[0].Addr.Addr, + Insecure: tc.InsecureSkipVerify, + }) + if err != nil { + return nil, trace.Wrap(err) + } + sshProxyHost, sshProxyPort, err := ping.Proxy.SSHProxyHostPort() + if err != nil { + return nil, trace.Wrap(err) + } + return &sshProxyParams{ + proxyHost: sshProxyHost, + proxyPort: sshProxyPort, + targetHost: targetHost, + targetPort: targetPort, + clusterName: ping.ClusterName, + tlsRouting: ping.Proxy.TLSRoutingEnabled, + }, nil +} + +// setupJumpHost configures the client for connecting to the jump host's proxy. +func setupJumpHost(cf *CLIConf, tc *libclient.TeleportClient, sp sshProxyParams) error { + return tc.WithoutJumpHosts(func(tc *libclient.TeleportClient) error { + // Fetch certificate for the leaf cluster. This allows users to log + // in once into the root cluster and let the proxy handle fetching + // certificates for leaf clusters automatically. + err := tc.LoadKeyForClusterWithReissue(cf.Context, sp.clusterName) + if err != nil { + return trace.Wrap(err) + } + // Update known_hosts with the leaf proxy's CA certificate, otherwise + // users will be prompted to manually accept the key. + err = tc.UpdateKnownHosts(cf.Context, sp.proxyHost, sp.clusterName) + if err != nil { + return trace.Wrap(err) + } + // We'll be connecting directly to the leaf cluster so make sure agent + // loads correct host CA. + tc.LocalAgent().UpdateCluster(sp.clusterName) + return nil + }) +} + +// sshProxyParams combines parameters for establishing an SSH proxy used +// as a ProxyCommand for SSH clients. +type sshProxyParams struct { + // proxyHost is the Teleport proxy host name. + proxyHost string + // proxyPort is the Teleport proxy port. + proxyPort string + // targetHost is the target SSH node host name. + targetHost string + // targetPort is the target SSH node port. + targetPort string + // clusterName is the cluster where the SSH node resides. + clusterName string + // tlsRouting is true if the Teleport proxy has TLS routing enabled. + tlsRouting bool } // cleanTargetHost cleans the targetHost and remote site and proxy suffixes. @@ -78,13 +171,8 @@ func cleanTargetHost(targetHost, proxyHost, siteName string) string { return targetHost } -func sshProxyWithTLSRouting(cf *CLIConf, tc *libclient.TeleportClient, targetHost, targetPort string) error { - address, err := utils.ParseAddr(tc.WebProxyAddr) - if err != nil { - return trace.Wrap(err) - } - - pool, err := tc.LocalAgent().ClientCertPool(tc.SiteName) +func sshProxyWithTLSRouting(cf *CLIConf, tc *libclient.TeleportClient, params sshProxyParams) error { + pool, err := tc.LocalAgent().ClientCertPool(params.clusterName) if err != nil { return trace.Wrap(err) } @@ -93,15 +181,15 @@ func sshProxyWithTLSRouting(cf *CLIConf, tc *libclient.TeleportClient, targetHos } lp, err := alpnproxy.NewLocalProxy(alpnproxy.LocalProxyConfig{ - RemoteProxyAddr: tc.WebProxyAddr, + RemoteProxyAddr: net.JoinHostPort(params.proxyHost, params.proxyPort), Protocols: []alpncommon.Protocol{alpncommon.ProtocolProxySSH}, InsecureSkipVerify: cf.InsecureSkipVerify, ParentContext: cf.Context, - SNI: address.Host(), + SNI: params.proxyHost, SSHUser: tc.HostLogin, - SSHUserHost: fmt.Sprintf("%s:%s", targetHost, targetPort), + SSHUserHost: fmt.Sprintf("%s:%s", params.targetHost, params.targetPort), SSHHostKeyCallback: tc.HostKeyCallback, - SSHTrustedCluster: cf.SiteName, + SSHTrustedCluster: params.clusterName, ClientTLSConfig: tlsConfig, }) if err != nil { @@ -114,7 +202,7 @@ func sshProxyWithTLSRouting(cf *CLIConf, tc *libclient.TeleportClient, targetHos return nil } -func sshProxy(tc *libclient.TeleportClient, targetHost, targetPort string) error { +func sshProxy(tc *libclient.TeleportClient, params sshProxyParams) error { sshPath, err := getSSHPath() if err != nil { return trace.Wrap(err) @@ -122,20 +210,21 @@ func sshProxy(tc *libclient.TeleportClient, targetHost, targetPort string) error keysDir := profile.FullProfilePath(tc.Config.KeysDir) knownHostsPath := keypaths.KnownHostsPath(keysDir) - sshHost, sshPort := tc.SSHProxyHostPort() args := []string{ "-A", "-o", fmt.Sprintf("UserKnownHostsFile=%s", knownHostsPath), - "-p", strconv.Itoa(sshPort), - sshHost, + "-p", params.proxyPort, + params.proxyHost, "-s", - fmt.Sprintf("proxy:%s:%s@%s", targetHost, targetPort, tc.SiteName), + fmt.Sprintf("proxy:%s:%s@%s", params.targetHost, params.targetPort, params.clusterName), } if tc.HostLogin != "" { args = append([]string{"-l", tc.HostLogin}, args...) } + log.Debugf("Executing proxy command: %v %v.", sshPath, strings.Join(args, " ")) + child := exec.Command(sshPath, args...) child.Stdin = os.Stdin child.Stdout = os.Stdout @@ -402,8 +491,8 @@ func onProxyCommandAWS(cf *CLIConf) error { } // loadAppCertificate loads the app certificate for the provided app. -func loadAppCertificate(tc *client.TeleportClient, appName string) (tls.Certificate, error) { - key, err := tc.LocalAgent().GetKey(tc.SiteName, client.WithAppCerts{}) +func loadAppCertificate(tc *libclient.TeleportClient, appName string) (tls.Certificate, error) { + key, err := tc.LocalAgent().GetKey(tc.SiteName, libclient.WithAppCerts{}) if err != nil { return tls.Certificate{}, trace.Wrap(err) } diff --git a/tool/tsh/proxy_test.go b/tool/tsh/proxy_test.go index cb7a48e680d..c32b524808e 100644 --- a/tool/tsh/proxy_test.go +++ b/tool/tsh/proxy_test.go @@ -337,6 +337,50 @@ func TestProxySSHDialWithIdentityFile(t *testing.T) { require.Contains(t, err.Error(), "subsystem request failed") } +// TestTSHProxyTemplate verifies connecting with OpenSSH client through the +// local proxy started with "tsh proxy ssh -J" using proxy templates. +func TestTSHProxyTemplate(t *testing.T) { + _, err := exec.LookPath("ssh") + if err != nil { + t.Skip("Skipping test, no ssh binary found.") + } + + lib.SetInsecureDevMode(true) + defer lib.SetInsecureDevMode(false) + + tshHome := t.TempDir() + t.Setenv(types.HomeEnvVar, tshHome) + + tshPath, err := os.Executable() + require.NoError(t, err) + + s := newTestSuite(t) + mustLogin(t, s) + + // Create proxy template configuration. + tshConfigFile := filepath.Join(tshHome, tshConfigPath) + require.NoError(t, os.MkdirAll(filepath.Dir(tshConfigFile), 0777)) + require.NoError(t, os.WriteFile(tshConfigFile, []byte(fmt.Sprintf(` +proxy_templates: +- template: '^(\w+)\.(root):(.+)$' + proxy: "%v" + host: "$1:$3" +`, s.root.Config.Proxy.WebAddr.Addr)), 0644)) + + // Create SSH config. + sshConfigFile := filepath.Join(tshHome, "sshconfig") + os.WriteFile(sshConfigFile, []byte(fmt.Sprintf(` +Host * + HostName %%h + StrictHostKeyChecking no + ProxyCommand %v -d --insecure proxy ssh -J {{proxy}} %%r@%%h:%%p +`, tshPath)), 0644) + + // Connect to "localnode" with OpenSSH. + mustRunOpenSSHCommand(t, sshConfigFile, "localnode.root", + s.root.Config.SSH.Addr.Port(defaults.SSHServerListenPort), "echo", "hello") +} + // TestTSHConfigConnectWithOpenSSHClient tests OpenSSH configuration generated by tsh config command and // connects to ssh node using native OpenSSH client with different session recording modes and proxy listener modes. func TestTSHConfigConnectWithOpenSSHClient(t *testing.T) { diff --git a/tool/tsh/tsh.go b/tool/tsh/tsh.go index e3ee894bdaa..cacbbcbdf7c 100644 --- a/tool/tsh/tsh.go +++ b/tool/tsh/tsh.go @@ -336,8 +336,8 @@ type CLIConf struct { // displayParticipantRequirements is set if verbose participant requirement information should be printed for moderated sessions. displayParticipantRequirements bool - // ExtraProxyHeaders is configuration read from the .tsh/config/config.yaml file. - ExtraProxyHeaders []ExtraProxyHeaders + // TshConfig is the loaded tsh configuration file ~/.tsh/config/config.yaml. + TshConfig TshConfig // SampleTraces indicates whether traces should be sampled. SampleTraces bool @@ -809,6 +809,7 @@ func Run(ctx context.Context, args []string, opts ...cliOption) error { if err != nil { return trace.Wrap(err) } + cf.TshConfig = *confOptions if cpuProfile != "" { log.Debugf("writing CPU profile to %v", cpuProfile) @@ -840,8 +841,6 @@ func Run(ctx context.Context, args []string, opts ...cliOption) error { }() } - cf.ExtraProxyHeaders = confOptions.ExtraHeaders - switch command { case ver.FullCommand(): err = onVersion(&cf) @@ -2477,7 +2476,7 @@ func makeClient(cf *CLIConf, useProfileLogin bool) (*client.TeleportClient, erro if strings.Contains(cf.UserHost, "=") { labels, err = client.ParseLabelSpec(cf.UserHost) if err != nil { - return nil, err + return nil, trace.Wrap(err) } } } else if cf.CopySpec != nil { @@ -2495,20 +2494,31 @@ func makeClient(cf *CLIConf, useProfileLogin bool) (*client.TeleportClient, erro } fPorts, err := client.ParsePortForwardSpec(cf.LocalForwardPorts) if err != nil { - return nil, err + return nil, trace.Wrap(err) } dPorts, err := client.ParseDynamicPortForwardSpec(cf.DynamicForwardedPorts) if err != nil { - return nil, err + return nil, trace.Wrap(err) } // 1: start with the defaults c := client.MakeDefaultConfig() + c.Host = cf.UserHost // ProxyJump is an alias of Proxy flag if cf.ProxyJump != "" { - hosts, err := utils.ParseProxyJump(cf.ProxyJump) + proxyJump := cf.ProxyJump + if strings.Contains(cf.ProxyJump, "{{proxy}}") { + proxy, host, matched := cf.TshConfig.ProxyTemplates.Apply(c.Host) + if !matched { + return nil, trace.BadParameter("proxy jump contains {{proxy}} variable but did not match any of the templates in tsh config") + } + proxyJump = strings.ReplaceAll(proxyJump, "{{proxy}}", proxy) + c.Host = host + log.Debugf("Will connect to proxy %q and host %q according to proxy templates.", proxyJump, host) + } + hosts, err := utils.ParseProxyJump(proxyJump) if err != nil { return nil, trace.Wrap(err) } @@ -2627,7 +2637,7 @@ func makeClient(cf *CLIConf, useProfileLogin bool) (*client.TeleportClient, erro if c.ExtraProxyHeaders == nil { c.ExtraProxyHeaders = map[string]string{} } - for _, proxyHeaders := range cf.ExtraProxyHeaders { + for _, proxyHeaders := range cf.TshConfig.ExtraHeaders { proxyGlob := utils.GlobToRegexp(proxyHeaders.Proxy) proxyRegexp, err := regexp.Compile(proxyGlob) if err != nil { @@ -2663,7 +2673,6 @@ func makeClient(cf *CLIConf, useProfileLogin bool) (*client.TeleportClient, erro if hostLogin != "" { c.HostLogin = hostLogin } - c.Host = cf.UserHost c.HostPort = int(cf.NodePort) c.Labels = labels c.KeyTTL = time.Minute * time.Duration(cf.MinsToLive) diff --git a/tool/tsh/tsh_helper_test.go b/tool/tsh/tsh_helper_test.go index 0afd3f5b7fa..a34c6bd972a 100644 --- a/tool/tsh/tsh_helper_test.go +++ b/tool/tsh/tsh_helper_test.go @@ -19,6 +19,7 @@ package main import ( "context" "fmt" + "net" "os/user" "testing" "time" @@ -40,6 +41,9 @@ type suite struct { } func (s *suite) setupRootCluster(t *testing.T, options testSuiteOptions) { + sshListenAddr := localListenerAddr() + _, sshListenPort, err := net.SplitHostPort(sshListenAddr) + require.NoError(t, err) fileConfig := &config.FileConfig{ Version: "v1", Global: config.Global{ @@ -55,10 +59,11 @@ func (s *suite) setupRootCluster(t *testing.T, options testSuiteOptions) { Proxy: config.Proxy{ Service: config.Service{ EnabledFlag: "true", - ListenAddress: localListenerAddr(), + ListenAddress: sshListenAddr, }, - WebAddr: localListenerAddr(), - TunAddr: localListenerAddr(), + SSHPublicAddr: []string{net.JoinHostPort("localhost", sshListenPort)}, + WebAddr: localListenerAddr(), + TunAddr: localListenerAddr(), }, Auth: config.Auth{ Service: config.Service{ @@ -71,7 +76,7 @@ func (s *suite) setupRootCluster(t *testing.T, options testSuiteOptions) { cfg := service.MakeDefaultConfig() cfg.CircuitBreakerConfig = breaker.NoopBreakerConfig() - err := config.ApplyFileConfig(fileConfig, cfg) + err = config.ApplyFileConfig(fileConfig, cfg) require.NoError(t, err) cfg.Proxy.DisableWebInterface = true diff --git a/tool/tsh/tsh_test.go b/tool/tsh/tsh_test.go index e60d27eb348..cc12079987e 100644 --- a/tool/tsh/tsh_test.go +++ b/tool/tsh/tsh_test.go @@ -435,7 +435,7 @@ func TestMakeClient(t *testing.T) { conf.NodePort = 46528 conf.LocalForwardPorts = []string{"80:remote:180"} conf.DynamicForwardedPorts = []string{":8080"} - conf.ExtraProxyHeaders = []ExtraProxyHeaders{ + conf.TshConfig.ExtraHeaders = []ExtraProxyHeaders{ {Proxy: "proxy:3080", Headers: map[string]string{"A": "B"}}, {Proxy: "*roxy:3080", Headers: map[string]string{"C": "D"}}, {Proxy: "*hello:3080", Headers: map[string]string{"E": "F"}}, // shouldn't get included diff --git a/tool/tsh/tshconfig.go b/tool/tsh/tshconfig.go index fa37d1df4f4..caa10949259 100644 --- a/tool/tsh/tshconfig.go +++ b/tool/tsh/tshconfig.go @@ -21,6 +21,8 @@ import ( "io/fs" "os" "path/filepath" + "regexp" + "strings" "github.com/gravitational/teleport/api/profile" @@ -41,6 +43,18 @@ type TshConfig struct { // ExtraHeaders are additional http headers to be included in // webclient requests. ExtraHeaders []ExtraProxyHeaders `yaml:"add_headers,omitempty"` + // ProxyTemplates describe rules for parsing out proxy out of full hostnames. + ProxyTemplates ProxyTemplates `yaml:"proxy_templates,omitempty"` +} + +// Check validates the tsh config. +func (config *TshConfig) Check() error { + for _, template := range config.ProxyTemplates { + if err := template.Check(); err != nil { + return trace.Wrap(err) + } + } + return nil } // ExtraProxyHeaders represents the headers to include with the @@ -64,13 +78,82 @@ func (config *TshConfig) Merge(otherConfig *TshConfig) TshConfig { } newConfig := TshConfig{} - - // extra headers - newConfig.ExtraHeaders = append(baseConfig.ExtraHeaders, otherConfig.ExtraHeaders...) + newConfig.ExtraHeaders = append(otherConfig.ExtraHeaders, baseConfig.ExtraHeaders...) + newConfig.ProxyTemplates = append(otherConfig.ProxyTemplates, baseConfig.ProxyTemplates...) return newConfig } +// ProxyTemplates represents a list of individual proxy templates. +type ProxyTemplates []*ProxyTemplate + +// Apply attempts to match the provided full hostname against all the templates +// in the list. Returns extracted proxy and host upon encountering the first +// matching template. +func (t ProxyTemplates) Apply(fullHostname string) (proxy, host string, matched bool) { + for _, template := range t { + proxy, host, matched := template.Apply(fullHostname) + if matched { + return proxy, host, true + } + } + return "", "", false +} + +// ProxyTemplate describes a single rule for parsing out proxy address from +// the full hostname. Used by tsh proxy ssh. +type ProxyTemplate struct { + // Template is a regular expression that full hostname is matched against. + Template string `yaml:"template"` + // Proxy is the proxy address. Can refer to regex groups from the template. + Proxy string `yaml:"proxy"` + // Host is optional hostname. Can refer to regex groups from the template. + Host string `yaml:"host"` + // re is the compiled template regexp. + re *regexp.Regexp +} + +// Check validates the proxy template. +func (t *ProxyTemplate) Check() (err error) { + if strings.TrimSpace(t.Proxy) == "" { + return trace.BadParameter("empty proxy expression") + } + if strings.TrimSpace(t.Template) == "" { + return trace.BadParameter("empty proxy template") + } + t.re, err = regexp.Compile(t.Template) + if err != nil { + return trace.Wrap(err) + } + return nil +} + +// Apply applies the proxy template to the provided hostname and returns +// expanded proxy address and hostname. +func (t ProxyTemplate) Apply(fullHostname string) (proxy, host string, matched bool) { + match := t.re.FindAllStringSubmatchIndex(fullHostname, -1) + if match == nil { + return "", "", false + } + + expandedProxy := []byte{} + for _, m := range match { + expandedProxy = t.re.ExpandString(expandedProxy, t.Proxy, fullHostname, m) + } + proxy = string(expandedProxy) + + host = fullHostname + if t.Host != "" { + expandedHost := []byte{} + for _, m := range match { + expandedHost = t.re.ExpandString(expandedHost, t.Host, fullHostname, m) + } + host = string(expandedHost) + } + + return proxy, host, true +} + // loadConfig load a single config file from given path. If the path does not exist, an empty config is returned instead. func loadConfig(fullConfigPath string) (*TshConfig, error) { bs, err := os.ReadFile(fullConfigPath) @@ -80,11 +163,13 @@ func loadConfig(fullConfigPath string) (*TshConfig, error) { } return nil, trace.ConvertSystemError(err) } - cfg := TshConfig{} if err := yaml.Unmarshal(bs, &cfg); err != nil { return nil, trace.ConvertSystemError(err) } + if err := cfg.Check(); err != nil { + return nil, trace.Wrap(err) + } return &cfg, nil } diff --git a/tool/tsh/tshconfig_test.go b/tool/tsh/tshconfig_test.go index 04c0eeb16b6..9a9eafba3c9 100644 --- a/tool/tsh/tshconfig_test.go +++ b/tool/tsh/tshconfig_test.go @@ -87,14 +87,14 @@ func TestLoadAllConfigs(t *testing.T) { require.NoError(t, err) require.Equal(t, &TshConfig{ ExtraHeaders: []ExtraProxyHeaders{ - { - Proxy: "global", - Headers: map[string]string{"bar": "123"}, - }, { Proxy: "user", Headers: map[string]string{"bar": "456"}, }, + { + Proxy: "global", + Headers: map[string]string{"bar": "123"}, + }, }, }, config) @@ -152,18 +152,18 @@ func TestTshConfigMerge(t *testing.T) { }}}, want: TshConfig{ ExtraHeaders: []ExtraProxyHeaders{ - { - Proxy: "foo", - Headers: map[string]string{ - "bar": "123", - }, - }, { Proxy: "bar", Headers: map[string]string{ "baz": "456", }, }, + { + Proxy: "foo", + Headers: map[string]string{ + "bar": "123", + }, + }, }}, }, { @@ -187,13 +187,13 @@ func TestTshConfigMerge(t *testing.T) { { Proxy: "foo", Headers: map[string]string{ - "bar": "123", + "bar": "456", }, }, { Proxy: "foo", Headers: map[string]string{ - "bar": "456", + "bar": "123", }, }, }}, @@ -207,3 +207,56 @@ func TestTshConfigMerge(t *testing.T) { }) } } + +// TestProxyTemplates verifies proxy templates matching functionality. +func TestProxyTemplates(t *testing.T) { + tshConfig := &TshConfig{ + ProxyTemplates: ProxyTemplates{ + { + Template: `^(.+)\.(us.example.com):(.+)$`, + Proxy: "$2:443", + Host: "$1:$3", + }, + { + Template: `^(.+)\.(eu.example.com):(.+)$`, + Proxy: "$2:3080", + }, + }, + } + require.NoError(t, tshConfig.Check()) + tests := []struct { + testName string + inFullHostname string + outProxy string + outHost string + outMatch bool + }{ + { + testName: "matches first template", + inFullHostname: "node-1.us.example.com:3022", + outProxy: "us.example.com:443", + outHost: "node-1:3022", + outMatch: true, + }, + { + testName: "matches second template", + inFullHostname: "node-1.eu.example.com:3022", + outProxy: "eu.example.com:3080", + outHost: "node-1.eu.example.com:3022", + outMatch: true, + }, + { + testName: "does not match templates", + inFullHostname: "node-1.cn.example.com:3022", + outMatch: false, + }, + } + for _, test := range tests { + t.Run(test.testName, func(t *testing.T) { + proxy, host, match := tshConfig.ProxyTemplates.Apply(test.inFullHostname) + require.Equal(t, test.outProxy, proxy) + require.Equal(t, test.outHost, host) + require.Equal(t, test.outMatch, match) + }) + } +}