mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-19 10:10:29 +08:00
Rename reversetunnelclient.RemoteSite (#57611)
This begins the process of cleaning up legacy terminology in lib/reversetunnel. The RemoteSite interface has been renamed Cluster to better reflect that it represents a Teleport cluster, either the root or leaf. The FakeRemoteSite has been renamed to FakeCluster, and all other mock and fake Cluster implementations have been renamed to use consitent terminology as well. All variables holding a reference to a cluster that were named site or had some reference to a site have also been changed to cluster or an equivalent which mentions cluster instead of site. An alias to RemoteSite has been left behind so that enterprise code can still consume it. It will be removed when the e ref is bumped at a later time when it is no longer referenced. Updates #19164.
This commit is contained in:
+1
-1
@@ -36,7 +36,7 @@ type ConnectionConfig struct {
|
||||
DesktopsGetter DesktopsGetter
|
||||
// Site represents a remote teleport site that can be accessed via
|
||||
// a teleport tunnel or directly by proxy.
|
||||
Site reversetunnelclient.RemoteSite
|
||||
Site reversetunnelclient.Cluster
|
||||
// ClientSrcAddr is the original observed client address.
|
||||
ClientSrcAddr net.Addr
|
||||
// ClientDstAddr is the original client's destination address.
|
||||
|
||||
@@ -319,9 +319,9 @@ func SetupTestContext(ctx context.Context, t *testing.T, cfg TestConfig) *TestCo
|
||||
testCtx.KubeProxy, err = proxy.NewTLSServer(proxy.TLSServerConfig{
|
||||
ForwarderConfig: proxy.ForwarderConfig{
|
||||
ReverseTunnelSrv: &reversetunnelclient.FakeServer{
|
||||
Sites: []reversetunnelclient.RemoteSite{
|
||||
&fakeRemoteSite{
|
||||
FakeRemoteSite: reversetunnelclient.NewFakeRemoteSite(testCtx.ClusterName, client),
|
||||
Clusters: []reversetunnelclient.Cluster{
|
||||
&fakeCluster{
|
||||
FakeCluster: reversetunnelclient.NewFakeCluster(testCtx.ClusterName, client),
|
||||
idToAddr: map[string]string{
|
||||
testCtx.HostID: testCtx.kubeServerListener.Addr().String(),
|
||||
},
|
||||
@@ -668,14 +668,14 @@ func (f *fakeClient) CreateSessionTracker(ctx context.Context, st types.SessionT
|
||||
}
|
||||
}
|
||||
|
||||
// fakeRemoteSite is a fake remote site that uses a map to map server IDs to
|
||||
// fakeCluster is a fake cluster that uses a map to map server IDs to
|
||||
// addresses to simulate reverse tunneling.
|
||||
type fakeRemoteSite struct {
|
||||
*reversetunnelclient.FakeRemoteSite
|
||||
type fakeCluster struct {
|
||||
*reversetunnelclient.FakeCluster
|
||||
idToAddr map[string]string
|
||||
}
|
||||
|
||||
func (f *fakeRemoteSite) DialTCP(p reversetunnelclient.DialParams) (conn net.Conn, err error) {
|
||||
func (f *fakeCluster) DialTCP(p reversetunnelclient.DialParams) (conn net.Conn, err error) {
|
||||
// The server ID is the first part of the address.
|
||||
addr, ok := f.idToAddr[strings.Split(p.ServerID, ".")[0]]
|
||||
if !ok {
|
||||
|
||||
@@ -125,9 +125,9 @@ func TestAuthenticate(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
tun := mockRevTunnel{
|
||||
sites: map[string]reversetunnelclient.RemoteSite{
|
||||
"remote": mockRemoteSite{name: "remote"},
|
||||
"local": mockRemoteSite{name: "local"},
|
||||
sites: map[string]reversetunnelclient.Cluster{
|
||||
"remote": mockCluster{name: "remote"},
|
||||
"local": mockCluster{name: "local"},
|
||||
},
|
||||
}
|
||||
f := &Forwarder{
|
||||
@@ -1159,15 +1159,15 @@ func (c *mockCAClient) GetCertAuthority(ctx context.Context, id types.CertAuthID
|
||||
return nil, trace.NotFound("cluster not found")
|
||||
}
|
||||
|
||||
// mockRemoteSite is a reversetunnelclient.RemoteSite implementation with hardcoded
|
||||
// mockCluster is a reversetunnelclient.Cluster implementation with hardcoded
|
||||
// name, because there's no easy way to construct a real
|
||||
// reversetunnelclient.RemoteSite.
|
||||
type mockRemoteSite struct {
|
||||
reversetunnelclient.RemoteSite
|
||||
// reversetunnelclient.Cluster.
|
||||
type mockCluster struct {
|
||||
reversetunnelclient.Cluster
|
||||
name string
|
||||
}
|
||||
|
||||
func (s mockRemoteSite) GetName() string { return s.name }
|
||||
func (s mockCluster) GetName() string { return s.name }
|
||||
|
||||
type mockAccessPoint struct {
|
||||
authclient.KubernetesAccessPoint
|
||||
@@ -1210,10 +1210,10 @@ func (ap mockAccessPoint) GetCertAuthority(ctx context.Context, id types.CertAut
|
||||
type mockRevTunnel struct {
|
||||
reversetunnelclient.Server
|
||||
|
||||
sites map[string]reversetunnelclient.RemoteSite
|
||||
sites map[string]reversetunnelclient.Cluster
|
||||
}
|
||||
|
||||
func (t mockRevTunnel) GetSite(name string) (reversetunnelclient.RemoteSite, error) {
|
||||
func (t mockRevTunnel) GetSite(name string) (reversetunnelclient.Cluster, error) {
|
||||
s, ok := t.sites[name]
|
||||
if !ok {
|
||||
return nil, trace.NotFound("remote site %q not found", name)
|
||||
@@ -1221,8 +1221,8 @@ func (t mockRevTunnel) GetSite(name string) (reversetunnelclient.RemoteSite, err
|
||||
return s, nil
|
||||
}
|
||||
|
||||
func (t mockRevTunnel) GetSites() ([]reversetunnelclient.RemoteSite, error) {
|
||||
var sites []reversetunnelclient.RemoteSite
|
||||
func (t mockRevTunnel) GetSites() ([]reversetunnelclient.Cluster, error) {
|
||||
var sites []reversetunnelclient.Cluster
|
||||
for _, s := range t.sites {
|
||||
sites = append(sites, s)
|
||||
}
|
||||
|
||||
@@ -110,7 +110,7 @@ type fakeReverseTunnel struct {
|
||||
t *testing.T
|
||||
}
|
||||
|
||||
func (f *fakeReverseTunnel) GetSite(_ string) (reversetunnelclient.RemoteSite, error) {
|
||||
func (f *fakeReverseTunnel) GetSite(_ string) (reversetunnelclient.Cluster, error) {
|
||||
return &fakeRemoteSiteTunnel{
|
||||
want: f.want,
|
||||
t: f.t,
|
||||
@@ -118,7 +118,7 @@ func (f *fakeReverseTunnel) GetSite(_ string) (reversetunnelclient.RemoteSite, e
|
||||
}
|
||||
|
||||
type fakeRemoteSiteTunnel struct {
|
||||
reversetunnelclient.RemoteSite
|
||||
reversetunnelclient.Cluster
|
||||
want reversetunnelclient.DialParams
|
||||
t *testing.T
|
||||
}
|
||||
|
||||
@@ -319,9 +319,9 @@ func SetupTestContext(ctx context.Context, t *testing.T, cfg TestConfig) *TestCo
|
||||
testCtx.KubeProxy, err = NewTLSServer(TLSServerConfig{
|
||||
ForwarderConfig: ForwarderConfig{
|
||||
ReverseTunnelSrv: &reversetunnelclient.FakeServer{
|
||||
Sites: []reversetunnelclient.RemoteSite{
|
||||
&fakeRemoteSite{
|
||||
FakeRemoteSite: reversetunnelclient.NewFakeRemoteSite(testCtx.ClusterName, client),
|
||||
Clusters: []reversetunnelclient.Cluster{
|
||||
&fakeCluster{
|
||||
FakeCluster: reversetunnelclient.NewFakeCluster(testCtx.ClusterName, client),
|
||||
idToAddr: map[string]string{
|
||||
testCtx.HostID: testCtx.kubeServerListener.Addr().String(),
|
||||
},
|
||||
@@ -688,14 +688,14 @@ func (f *fakeClient) CreateSessionTracker(ctx context.Context, st types.SessionT
|
||||
}
|
||||
}
|
||||
|
||||
// fakeRemoteSite is a fake remote site that uses a map to map server IDs to
|
||||
// fakeCluster is a fake cluster that uses a map to map server IDs to
|
||||
// addresses to simulate reverse tunneling.
|
||||
type fakeRemoteSite struct {
|
||||
*reversetunnelclient.FakeRemoteSite
|
||||
type fakeCluster struct {
|
||||
*reversetunnelclient.FakeCluster
|
||||
idToAddr map[string]string
|
||||
}
|
||||
|
||||
func (f *fakeRemoteSite) DialTCP(p reversetunnelclient.DialParams) (conn net.Conn, err error) {
|
||||
func (f *fakeCluster) DialTCP(p reversetunnelclient.DialParams) (conn net.Conn, err error) {
|
||||
// The server ID is the first part of the address.
|
||||
addr, ok := f.idToAddr[strings.Split(p.ServerID, ".")[0]]
|
||||
if !ok {
|
||||
|
||||
+57
-57
@@ -103,13 +103,13 @@ func (c *ProxiedMetricConn) Close() error {
|
||||
return trace.Wrap(c.Conn.Close())
|
||||
}
|
||||
|
||||
type serverResolverFn = func(ctx context.Context, host, port string, site site) (types.Server, error)
|
||||
type serverResolverFn = func(ctx context.Context, host, port string, cluster cluster) (types.Server, error)
|
||||
type windowsDesktopServiceConnectorFn = func(ctx context.Context, config *desktop.ConnectionConfig) (conn net.Conn, version string, err error)
|
||||
|
||||
// SiteGetter provides access to connected local or remote sites
|
||||
// SiteGetter provides access to connected local or remote clusters.
|
||||
type SiteGetter interface {
|
||||
// GetSite returns the site matching the provided clusterName
|
||||
GetSite(clusterName string) (reversetunnelclient.RemoteSite, error)
|
||||
// GetSite returns the cluster matching the provided clusterName
|
||||
GetSite(clusterName string) (reversetunnelclient.Cluster, error)
|
||||
}
|
||||
|
||||
// LocalAccessPoint provides access to remote cluster resources
|
||||
@@ -127,7 +127,7 @@ type RouterConfig struct {
|
||||
ClusterName string
|
||||
// LocalAccessPoint is the proxy cache
|
||||
LocalAccessPoint LocalAccessPoint
|
||||
// SiteGetter allows looking up sites
|
||||
// SiteGetter allows looking up clusters
|
||||
SiteGetter SiteGetter
|
||||
// TracerProvider allows tracers to be created
|
||||
TracerProvider oteltrace.TracerProvider
|
||||
@@ -178,7 +178,7 @@ func (c *RouterConfig) CheckAndSetDefaults() error {
|
||||
type Router struct {
|
||||
clusterName string
|
||||
localAccessPoint LocalAccessPoint
|
||||
localSite reversetunnelclient.RemoteSite
|
||||
localCluster reversetunnelclient.Cluster
|
||||
siteGetter SiteGetter
|
||||
tracer oteltrace.Tracer
|
||||
log *slog.Logger
|
||||
@@ -193,7 +193,7 @@ func NewRouter(cfg RouterConfig) (*Router, error) {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
localSite, err := cfg.SiteGetter.GetSite(cfg.ClusterName)
|
||||
localCluster, err := cfg.SiteGetter.GetSite(cfg.ClusterName)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -201,7 +201,7 @@ func NewRouter(cfg RouterConfig) (*Router, error) {
|
||||
return &Router{
|
||||
clusterName: cfg.ClusterName,
|
||||
localAccessPoint: cfg.LocalAccessPoint,
|
||||
localSite: localSite,
|
||||
localCluster: localCluster,
|
||||
siteGetter: cfg.SiteGetter,
|
||||
tracer: cfg.TracerProvider.Tracer("Router"),
|
||||
log: cfg.Logger,
|
||||
@@ -231,17 +231,17 @@ func (r *Router) DialHost(ctx context.Context, clientSrcAddr, clientDstAddr net.
|
||||
tracing.EndSpan(span, err)
|
||||
}()
|
||||
|
||||
site := r.localSite
|
||||
cluster := r.localCluster
|
||||
if clusterName != r.clusterName {
|
||||
remoteSite, err := r.getRemoteCluster(ctx, clusterName, clusterAccessChecker)
|
||||
remoteCluster, err := r.getRemoteCluster(ctx, clusterName, clusterAccessChecker)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "looking up remote cluster %q", clusterName)
|
||||
}
|
||||
site = remoteSite
|
||||
cluster = remoteCluster
|
||||
}
|
||||
|
||||
span.AddEvent("looking up server")
|
||||
target, err := r.serverResolver(ctx, host, port, remoteSite{site})
|
||||
target, err := r.serverResolver(ctx, host, port, fakeCluster{cluster})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -304,7 +304,7 @@ func (r *Router) DialHost(ctx context.Context, clientSrcAddr, clientDstAddr net.
|
||||
return nil, trace.ConnectionProblem(errors.New("connection problem"), "direct dialing to nodes not found in inventory is not supported")
|
||||
}
|
||||
|
||||
conn, err := site.Dial(reversetunnelclient.DialParams{
|
||||
conn, err := cluster.Dial(reversetunnelclient.DialParams{
|
||||
From: clientSrcAddr,
|
||||
To: &utils.NetAddr{AddrNetwork: "tcp", Addr: serverAddr},
|
||||
OriginalClientDstAddr: clientDstAddr,
|
||||
@@ -340,16 +340,16 @@ func (r *Router) DialWindowsDesktop(ctx context.Context, clientSrcAddr, clientDs
|
||||
)
|
||||
defer func() { tracing.EndSpan(span, err) }()
|
||||
|
||||
site := r.localSite
|
||||
cluster := r.localCluster
|
||||
if clusterName != r.clusterName {
|
||||
remoteSite, err := r.getRemoteCluster(ctx, clusterName, clusterAccessChecker)
|
||||
remoteCluster, err := r.getRemoteCluster(ctx, clusterName, clusterAccessChecker)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err, "looking up remote cluster %q", clusterName)
|
||||
}
|
||||
site = remoteSite
|
||||
cluster = remoteCluster
|
||||
}
|
||||
|
||||
accessPoint, err := site.CachingAccessPoint()
|
||||
accessPoint, err := cluster.CachingAccessPoint()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -359,7 +359,7 @@ func (r *Router) DialWindowsDesktop(ctx context.Context, clientSrcAddr, clientDs
|
||||
serviceConn, _, err := r.windowsDesktopServiceConnector(ctx, &desktop.ConnectionConfig{
|
||||
Log: r.log,
|
||||
DesktopsGetter: accessPoint,
|
||||
Site: site,
|
||||
Site: cluster,
|
||||
ClientSrcAddr: clientSrcAddr,
|
||||
ClientDstAddr: clientDstAddr,
|
||||
ClusterName: clusterName,
|
||||
@@ -412,9 +412,9 @@ func (c *checkedPrefixWriter) Write(p []byte) (int, error) {
|
||||
return n, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// getRemoteCluster looks up the provided clusterName to determine if a remote site exists with
|
||||
// getRemoteCluster looks up the provided clusterName to determine if a remote cluster exists with
|
||||
// that name and determines if the user has access to it.
|
||||
func (r *Router) getRemoteCluster(ctx context.Context, clusterName string, clusterAccessChecker func(types.RemoteCluster) error) (reversetunnelclient.RemoteSite, error) {
|
||||
func (r *Router) getRemoteCluster(ctx context.Context, clusterName string, clusterAccessChecker func(types.RemoteCluster) error) (reversetunnelclient.Cluster, error) {
|
||||
_, span := r.tracer.Start(
|
||||
ctx,
|
||||
"router/getRemoteCluster",
|
||||
@@ -424,7 +424,7 @@ func (r *Router) getRemoteCluster(ctx context.Context, clusterName string, clust
|
||||
)
|
||||
defer span.End()
|
||||
|
||||
site, err := r.siteGetter.GetSite(clusterName)
|
||||
cluster, err := r.siteGetter.GetSite(clusterName)
|
||||
if err != nil {
|
||||
return nil, utils.OpaqueAccessDenied(err)
|
||||
}
|
||||
@@ -438,26 +438,26 @@ func (r *Router) getRemoteCluster(ctx context.Context, clusterName string, clust
|
||||
return nil, utils.OpaqueAccessDenied(err)
|
||||
}
|
||||
|
||||
return site, nil
|
||||
return cluster, nil
|
||||
}
|
||||
|
||||
// site is the minimum interface needed to match servers
|
||||
// for a reversetunnelclient.RemoteSite. It makes testing easier.
|
||||
type site interface {
|
||||
// cluster is the minimum interface needed to match servers
|
||||
// for a reversetunnelclient.Cluster. It makes testing easier.
|
||||
type cluster interface {
|
||||
GetNodes(ctx context.Context, fn func(n readonly.Server) bool) ([]types.Server, error)
|
||||
GetClusterNetworkingConfig(ctx context.Context) (types.ClusterNetworkingConfig, error)
|
||||
GetGitServers(context.Context, func(readonly.Server) bool) ([]types.Server, error)
|
||||
}
|
||||
|
||||
// remoteSite is a site implementation that wraps
|
||||
// a reversetunnelclient.RemoteSite
|
||||
type remoteSite struct {
|
||||
site reversetunnelclient.RemoteSite
|
||||
// fakeCluster is a cluster implementation that wraps
|
||||
// a reversetunnelclient.Cluster
|
||||
type fakeCluster struct {
|
||||
cluster reversetunnelclient.Cluster
|
||||
}
|
||||
|
||||
// GetNodes uses the wrapped sites NodeWatcher to filter nodes
|
||||
func (r remoteSite) GetNodes(ctx context.Context, fn func(n readonly.Server) bool) ([]types.Server, error) {
|
||||
watcher, err := r.site.NodeWatcher()
|
||||
// GetNodes uses the wrapped cluster's NodeWatcher to filter nodes
|
||||
func (r fakeCluster) GetNodes(ctx context.Context, fn func(n readonly.Server) bool) ([]types.Server, error) {
|
||||
watcher, err := r.cluster.NodeWatcher()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -466,9 +466,9 @@ func (r remoteSite) GetNodes(ctx context.Context, fn func(n readonly.Server) boo
|
||||
return servers, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// GetGitServers uses the wrapped sites GitServerWatcher to filter git servers.
|
||||
func (r remoteSite) GetGitServers(ctx context.Context, fn func(n readonly.Server) bool) ([]types.Server, error) {
|
||||
watcher, err := r.site.GitServerWatcher()
|
||||
// GetGitServers uses the wrapped cluster's GitServerWatcher to filter git servers.
|
||||
func (r fakeCluster) GetGitServers(ctx context.Context, fn func(n readonly.Server) bool) ([]types.Server, error) {
|
||||
watcher, err := r.cluster.GitServerWatcher()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -476,9 +476,9 @@ func (r remoteSite) GetGitServers(ctx context.Context, fn func(n readonly.Server
|
||||
return watcher.CurrentResourcesWithFilter(ctx, fn)
|
||||
}
|
||||
|
||||
// GetClusterNetworkingConfig uses the wrapped sites cache to retrieve the ClusterNetworkingConfig
|
||||
func (r remoteSite) GetClusterNetworkingConfig(ctx context.Context) (types.ClusterNetworkingConfig, error) {
|
||||
ap, err := r.site.CachingAccessPoint()
|
||||
// GetClusterNetworkingConfig uses the wrapped cluster's cache to retrieve the ClusterNetworkingConfig
|
||||
func (r fakeCluster) GetClusterNetworkingConfig(ctx context.Context) (types.ClusterNetworkingConfig, error) {
|
||||
ap, err := r.cluster.CachingAccessPoint()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -488,27 +488,27 @@ func (r remoteSite) GetClusterNetworkingConfig(ctx context.Context) (types.Clust
|
||||
}
|
||||
|
||||
// getServer attempts to locate a node matching the provided host and port in
|
||||
// the provided site.
|
||||
func getServer(ctx context.Context, host, port string, site site) (types.Server, error) {
|
||||
// the provided cluster.
|
||||
func getServer(ctx context.Context, host, port string, cluster cluster) (types.Server, error) {
|
||||
if org, ok := types.GetGitHubOrgFromNodeAddr(host); ok {
|
||||
return getGitHubServer(ctx, org, site)
|
||||
return getGitHubServer(ctx, org, cluster)
|
||||
}
|
||||
return getServerWithResolver(ctx, host, port, site, nil /* use default resolver */)
|
||||
return getServerWithResolver(ctx, host, port, cluster, nil /* use default resolver */)
|
||||
}
|
||||
|
||||
var disableUnqualifiedLookups = os.Getenv("TELEPORT_UNSTABLE_DISABLE_UNQUALIFIED_LOOKUPS") == "yes"
|
||||
|
||||
// getServerWithResolver attempts to locate a node matching the provided host and port in
|
||||
// the provided site. The resolver argument is used in certain tests to mock DNS resolution
|
||||
// the provided cluster. The resolver argument is used in certain tests to mock DNS resolution
|
||||
// and can generally be left nil.
|
||||
func getServerWithResolver(ctx context.Context, host, port string, site site, resolver apiutils.HostResolver) (types.Server, error) {
|
||||
if site == nil {
|
||||
return nil, trace.BadParameter("invalid remote site provided")
|
||||
func getServerWithResolver(ctx context.Context, host, port string, cluster cluster, resolver apiutils.HostResolver) (types.Server, error) {
|
||||
if cluster == nil {
|
||||
return nil, trace.BadParameter("invalid remote cluster provided")
|
||||
}
|
||||
|
||||
strategy := types.RoutingStrategy_UNAMBIGUOUS_MATCH
|
||||
var caseInsensitiveRouting bool
|
||||
if cfg, err := site.GetClusterNetworkingConfig(ctx); err == nil {
|
||||
if cfg, err := cluster.GetClusterNetworkingConfig(ctx); err == nil {
|
||||
strategy = cfg.GetRoutingStrategy()
|
||||
caseInsensitiveRouting = cfg.GetCaseInsensitiveRouting()
|
||||
}
|
||||
@@ -526,7 +526,7 @@ func getServerWithResolver(ctx context.Context, host, port string, site site, re
|
||||
|
||||
var maxScore int
|
||||
scores := make(map[string]int)
|
||||
matches, err := site.GetNodes(ctx, func(server readonly.Server) bool {
|
||||
matches, err := cluster.GetNodes(ctx, func(server readonly.Server) bool {
|
||||
score := routeMatcher.RouteToServerScore(server)
|
||||
if score < 1 {
|
||||
return false
|
||||
@@ -616,17 +616,17 @@ func (r *Router) DialSite(ctx context.Context, clusterName string, clientSrcAddr
|
||||
|
||||
// dial the local auth server
|
||||
if clusterName == r.clusterName {
|
||||
conn, err := r.localSite.DialAuthServer(reversetunnelclient.DialParams{From: clientSrcAddr, OriginalClientDstAddr: clientDstAddr})
|
||||
conn, err := r.localCluster.DialAuthServer(reversetunnelclient.DialParams{From: clientSrcAddr, OriginalClientDstAddr: clientDstAddr})
|
||||
return conn, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// lookup the site and dial its auth server
|
||||
site, err := r.siteGetter.GetSite(clusterName)
|
||||
// lookup the cluster and dial its auth server
|
||||
cluster, err := r.siteGetter.GetSite(clusterName)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
conn, err := site.DialAuthServer(reversetunnelclient.DialParams{From: clientSrcAddr, OriginalClientDstAddr: clientDstAddr})
|
||||
conn, err := cluster.DialAuthServer(reversetunnelclient.DialParams{From: clientSrcAddr, OriginalClientDstAddr: clientDstAddr})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -637,18 +637,18 @@ func (r *Router) DialSite(ctx context.Context, clusterName string, clientSrcAddr
|
||||
// GetSiteClient returns an auth client for the provided cluster.
|
||||
func (r *Router) GetSiteClient(ctx context.Context, clusterName string) (authclient.ClientI, error) {
|
||||
if clusterName == r.clusterName {
|
||||
return r.localSite.GetClient()
|
||||
return r.localCluster.GetClient()
|
||||
}
|
||||
|
||||
site, err := r.siteGetter.GetSite(clusterName)
|
||||
cluster, err := r.siteGetter.GetSite(clusterName)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
return site.GetClient()
|
||||
return cluster.GetClient()
|
||||
}
|
||||
|
||||
func getGitHubServer(ctx context.Context, gitHubOrg string, site site) (types.Server, error) {
|
||||
servers, err := site.GetGitServers(ctx, func(s readonly.Server) bool {
|
||||
func getGitHubServer(ctx context.Context, gitHubOrg string, cluster cluster) (types.Server, error) {
|
||||
servers, err := cluster.GetGitServers(ctx, func(s readonly.Server) bool {
|
||||
github := s.GetGitHub()
|
||||
return github != nil && github.Organization == gitHubOrg
|
||||
})
|
||||
|
||||
+37
-37
@@ -541,7 +541,7 @@ func TestGetServers(t *testing.T) {
|
||||
}
|
||||
|
||||
func serverResolver(srv types.Server, err error) serverResolverFn {
|
||||
return func(ctx context.Context, host, port string, site site) (types.Server, error) {
|
||||
return func(ctx context.Context, host, port string, site cluster) (types.Server, error) {
|
||||
return srv, err
|
||||
}
|
||||
}
|
||||
@@ -627,19 +627,19 @@ func TestCheckedPrefixWriter(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
type tunnel struct {
|
||||
type fakeTunnel struct {
|
||||
reversetunnelclient.Tunnel
|
||||
|
||||
site reversetunnelclient.RemoteSite
|
||||
err error
|
||||
cluster reversetunnelclient.Cluster
|
||||
err error
|
||||
}
|
||||
|
||||
func (t tunnel) GetSite(cluster string) (reversetunnelclient.RemoteSite, error) {
|
||||
return t.site, t.err
|
||||
func (t fakeTunnel) GetSite(cluster string) (reversetunnelclient.Cluster, error) {
|
||||
return t.cluster, t.err
|
||||
}
|
||||
|
||||
type testRemoteSite struct {
|
||||
reversetunnelclient.RemoteSite
|
||||
reversetunnelclient.Cluster
|
||||
|
||||
params reversetunnelclient.DialParams
|
||||
|
||||
@@ -664,12 +664,12 @@ func (r testRemoteSite) CachingAccessPoint() (authclient.RemoteProxyAccessPoint,
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
type testSiteGetter struct {
|
||||
site reversetunnelclient.RemoteSite
|
||||
type fakeSiteGetter struct {
|
||||
cluster reversetunnelclient.Cluster
|
||||
}
|
||||
|
||||
func (s testSiteGetter) GetSite(clusterName string) (reversetunnelclient.RemoteSite, error) {
|
||||
return s.site, nil
|
||||
func (s fakeSiteGetter) GetSite(clusterName string) (reversetunnelclient.Cluster, error) {
|
||||
return s.cluster, nil
|
||||
}
|
||||
|
||||
type fakeConn struct {
|
||||
@@ -748,7 +748,7 @@ func TestRouter_DialHost(t *testing.T) {
|
||||
name: "failure looking up cluster",
|
||||
router: Router{
|
||||
clusterName: "leaf",
|
||||
siteGetter: tunnel{err: trace.NotFound("unknown cluster")},
|
||||
siteGetter: fakeTunnel{err: trace.NotFound("unknown cluster")},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
},
|
||||
assertion: func(t *testing.T, params reversetunnelclient.DialParams, conn net.Conn, err error) {
|
||||
@@ -761,7 +761,7 @@ func TestRouter_DialHost(t *testing.T) {
|
||||
name: "dial failure",
|
||||
router: Router{
|
||||
clusterName: "test",
|
||||
localSite: &testRemoteSite{err: trace.ConnectionProblem(context.DeadlineExceeded, "connection refused")},
|
||||
localCluster: &testRemoteSite{err: trace.ConnectionProblem(context.DeadlineExceeded, "connection refused")},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
serverResolver: serverResolver(srv, nil),
|
||||
},
|
||||
@@ -775,7 +775,7 @@ func TestRouter_DialHost(t *testing.T) {
|
||||
name: "dial success",
|
||||
router: Router{
|
||||
clusterName: "test",
|
||||
localSite: &testRemoteSite{conn: fakeConn{}},
|
||||
localCluster: &testRemoteSite{conn: fakeConn{}},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
serverResolver: serverResolver(srv, nil),
|
||||
},
|
||||
@@ -793,8 +793,8 @@ func TestRouter_DialHost(t *testing.T) {
|
||||
name: "dial success to agentless node",
|
||||
router: Router{
|
||||
clusterName: "test",
|
||||
localSite: &testRemoteSite{conn: fakeConn{}},
|
||||
siteGetter: &testSiteGetter{site: &testRemoteSite{conn: fakeConn{}}},
|
||||
localCluster: &testRemoteSite{conn: fakeConn{}},
|
||||
siteGetter: &fakeSiteGetter{cluster: &testRemoteSite{conn: fakeConn{}}},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
serverResolver: serverResolver(agentlessSrv, nil),
|
||||
},
|
||||
@@ -813,8 +813,8 @@ func TestRouter_DialHost(t *testing.T) {
|
||||
name: "dial success to agentless node using EC2 Instance Connect Endpoint",
|
||||
router: Router{
|
||||
clusterName: "test",
|
||||
localSite: &testRemoteSite{conn: fakeConn{}},
|
||||
siteGetter: &testSiteGetter{site: &testRemoteSite{conn: fakeConn{}}},
|
||||
localCluster: &testRemoteSite{conn: fakeConn{}},
|
||||
siteGetter: &fakeSiteGetter{cluster: &testRemoteSite{conn: fakeConn{}}},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
serverResolver: serverResolver(agentlessEC2ICESrv, nil),
|
||||
},
|
||||
@@ -836,8 +836,8 @@ func TestRouter_DialHost(t *testing.T) {
|
||||
conn, err := tt.router.DialHost(ctx, &utils.NetAddr{}, &utils.NetAddr{}, "host", "0", "test", nil, agentGetter, createSigner)
|
||||
|
||||
var params reversetunnelclient.DialParams
|
||||
if tt.router.localSite != nil {
|
||||
params = tt.router.localSite.(*testRemoteSite).params
|
||||
if tt.router.localCluster != nil {
|
||||
params = tt.router.localCluster.(*testRemoteSite).params
|
||||
}
|
||||
|
||||
tt.assertion(t, params, conn, err)
|
||||
@@ -854,7 +854,7 @@ func TestRouter_DialSite(t *testing.T) {
|
||||
name string
|
||||
cluster string
|
||||
localSite testRemoteSite
|
||||
tunnel tunnel
|
||||
tunnel fakeTunnel
|
||||
assertion func(t *testing.T, conn net.Conn, err error)
|
||||
}{
|
||||
{
|
||||
@@ -888,8 +888,8 @@ func TestRouter_DialSite(t *testing.T) {
|
||||
{
|
||||
name: "failure to dial remote site",
|
||||
cluster: "leaf",
|
||||
tunnel: tunnel{
|
||||
site: &testRemoteSite{err: trace.ConnectionProblem(context.DeadlineExceeded, "connection refused")},
|
||||
tunnel: fakeTunnel{
|
||||
cluster: &testRemoteSite{err: trace.ConnectionProblem(context.DeadlineExceeded, "connection refused")},
|
||||
},
|
||||
assertion: func(t *testing.T, conn net.Conn, err error) {
|
||||
require.Error(t, err)
|
||||
@@ -900,7 +900,7 @@ func TestRouter_DialSite(t *testing.T) {
|
||||
{
|
||||
name: "unknown cluster",
|
||||
cluster: "fake",
|
||||
tunnel: tunnel{
|
||||
tunnel: fakeTunnel{
|
||||
err: trace.NotFound("unknown cluster"),
|
||||
},
|
||||
assertion: func(t *testing.T, conn net.Conn, err error) {
|
||||
@@ -912,8 +912,8 @@ func TestRouter_DialSite(t *testing.T) {
|
||||
{
|
||||
name: "successfully dial remote site",
|
||||
cluster: "leaf",
|
||||
tunnel: tunnel{
|
||||
site: &testRemoteSite{conn: fakeConn{}},
|
||||
tunnel: fakeTunnel{
|
||||
cluster: &testRemoteSite{conn: fakeConn{}},
|
||||
},
|
||||
assertion: func(t *testing.T, conn net.Conn, err error) {
|
||||
require.NoError(t, err)
|
||||
@@ -927,10 +927,10 @@ func TestRouter_DialSite(t *testing.T) {
|
||||
for _, tt := range cases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
router := Router{
|
||||
clusterName: cluster,
|
||||
localSite: &tt.localSite,
|
||||
siteGetter: tt.tunnel,
|
||||
tracer: tracing.NoopTracer(cluster),
|
||||
clusterName: cluster,
|
||||
localCluster: &tt.localSite,
|
||||
siteGetter: tt.tunnel,
|
||||
tracer: tracing.NoopTracer(cluster),
|
||||
}
|
||||
|
||||
conn, err := router.DialSite(ctx, tt.cluster, nil, nil)
|
||||
@@ -951,7 +951,7 @@ func TestRouter_DialWindowsDesktop(t *testing.T) {
|
||||
name: "failure looking up cluster",
|
||||
router: Router{
|
||||
clusterName: "leaf",
|
||||
siteGetter: tunnel{err: trace.NotFound("unknown cluster")},
|
||||
siteGetter: fakeTunnel{err: trace.NotFound("unknown cluster")},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
},
|
||||
assertion: func(t *testing.T, conn net.Conn, err error) {
|
||||
@@ -963,9 +963,9 @@ func TestRouter_DialWindowsDesktop(t *testing.T) {
|
||||
{
|
||||
name: "failure connecting to desktop service",
|
||||
router: Router{
|
||||
clusterName: "test",
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
localSite: &testRemoteSite{},
|
||||
clusterName: "test",
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
localCluster: &testRemoteSite{},
|
||||
windowsDesktopServiceConnector: func(ctx context.Context, c *desktop.ConnectionConfig) (net.Conn, string, error) {
|
||||
return nil, "", trace.ConnectionProblem(context.DeadlineExceeded, "connection refused")
|
||||
},
|
||||
@@ -979,9 +979,9 @@ func TestRouter_DialWindowsDesktop(t *testing.T) {
|
||||
{
|
||||
name: "dial success",
|
||||
router: Router{
|
||||
clusterName: "test",
|
||||
localSite: &testRemoteSite{conn: fakeConn{}},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
clusterName: "test",
|
||||
localCluster: &testRemoteSite{conn: fakeConn{}},
|
||||
tracer: tracing.NoopTracer("test"),
|
||||
windowsDesktopServiceConnector: func(ctx context.Context, c *desktop.ConnectionConfig) (net.Conn, string, error) {
|
||||
return fakeConn{}, "18.0.0", nil
|
||||
},
|
||||
|
||||
@@ -563,7 +563,7 @@ func (s *remoteSite) updateCertAuthorities(retry retryutils.Retry, remoteWatcher
|
||||
s.logger.DebugContext(s.ctx, "Remote cluster does not support cert authorities rotation yet")
|
||||
case trace.IsCompareFailed(err):
|
||||
s.logger.InfoContext(s.ctx, "Remote cluster has updated certificate authorities, going to force reconnect")
|
||||
if err := s.srv.onSiteTunnelClose(&alwaysClose{RemoteSite: s}); err != nil {
|
||||
if err := s.srv.onSiteTunnelClose(&alwaysClose{Cluster: s}); err != nil {
|
||||
s.logger.WarnContext(s.ctx, "Failed to close remote site", "error", err)
|
||||
}
|
||||
return
|
||||
|
||||
@@ -417,7 +417,7 @@ func (s *server) disconnectClusters(connectedRemoteClusters []*remoteSite, remot
|
||||
for _, cluster := range connectedRemoteClusters {
|
||||
if _, ok := remoteMap[cluster.GetName()]; !ok {
|
||||
s.logger.InfoContext(s.ctx, "Remote cluster has been deleted, disconnecting it from the proxy", "remote_cluster", cluster.GetName())
|
||||
if err := s.onSiteTunnelClose(&alwaysClose{RemoteSite: cluster}); err != nil {
|
||||
if err := s.onSiteTunnelClose(&alwaysClose{Cluster: cluster}); err != nil {
|
||||
s.logger.DebugContext(s.ctx, "Failure closing cluster", "remote_cluster", cluster.GetName(), "error", err)
|
||||
}
|
||||
remoteClustersStats.DeleteLabelValues(cluster.GetName())
|
||||
@@ -1067,10 +1067,10 @@ func (s *server) upsertRemoteCluster(conn net.Conn, sshConn *ssh.ServerConn) (*r
|
||||
return site, remoteConn, nil
|
||||
}
|
||||
|
||||
func (s *server) GetSites() ([]reversetunnelclient.RemoteSite, error) {
|
||||
func (s *server) GetSites() ([]reversetunnelclient.Cluster, error) {
|
||||
s.RLock()
|
||||
defer s.RUnlock()
|
||||
out := make([]reversetunnelclient.RemoteSite, 0, len(s.remoteSites)+len(s.clusterPeers)+1)
|
||||
out := make([]reversetunnelclient.Cluster, 0, len(s.remoteSites)+len(s.clusterPeers)+1)
|
||||
out = append(out, s.localSite)
|
||||
|
||||
haveLocalConnection := make(map[string]bool)
|
||||
@@ -1104,7 +1104,7 @@ func (s *server) getRemoteClusters() []*remoteSite {
|
||||
// with a cluster peer your best bet is to wait until the agent has discovered
|
||||
// all proxies behind a load balancer. Note, the cluster peer is a
|
||||
// services.TunnelConnection that was created by another proxy.
|
||||
func (s *server) GetSite(name string) (reversetunnelclient.RemoteSite, error) {
|
||||
func (s *server) GetSite(name string) (reversetunnelclient.Cluster, error) {
|
||||
s.RLock()
|
||||
defer s.RUnlock()
|
||||
if s.localSite.GetName() == name {
|
||||
@@ -1131,7 +1131,7 @@ func (s *server) GetProxyPeerClient() *peer.Client {
|
||||
// alwaysClose forces onSiteTunnelClose to remove and close
|
||||
// the site by always returning false from HasValidConnections.
|
||||
type alwaysClose struct {
|
||||
reversetunnelclient.RemoteSite
|
||||
reversetunnelclient.Cluster
|
||||
}
|
||||
|
||||
func (a *alwaysClose) HasValidConnections() bool {
|
||||
|
||||
@@ -105,56 +105,59 @@ func (params DialParams) String() string {
|
||||
return fmt.Sprintf("from: %q to: %q", params.From, to)
|
||||
}
|
||||
|
||||
// RemoteSite represents remote teleport site that can be accessed via
|
||||
// teleport tunnel or directly by proxy
|
||||
//
|
||||
// There are two implementations of this interface: local and remote sites.
|
||||
type RemoteSite interface {
|
||||
// DialAuthServer returns a net.Conn to the Auth Server of a site.
|
||||
// RemoteSite is an alis to allow migrating to Cluster without breaking builds.
|
||||
// Deprecated: Use Cluster instead
|
||||
// TODO(tross): Delete when all references are converted to Cluster.
|
||||
type RemoteSite = Cluster
|
||||
|
||||
// Cluster represents a teleport cluster, either root or leaf,
|
||||
// that can be accessed via teleport tunnel or directly by proxy.
|
||||
type Cluster interface {
|
||||
// DialAuthServer returns a net.Conn to the Auth Server of a cluster.
|
||||
DialAuthServer(DialParams) (conn net.Conn, err error)
|
||||
// Dial dials any address within the site network, in terminating
|
||||
// Dial dials any address within the cluster network, in terminating
|
||||
// mode it uses local instance of forwarding server to terminate
|
||||
// and record the connection.
|
||||
Dial(DialParams) (conn net.Conn, err error)
|
||||
// DialTCP dials any address within the site network and
|
||||
// DialTCP dials any address within the cluster network and
|
||||
// ignores recording mode, used in components that need direct dialer.
|
||||
DialTCP(DialParams) (conn net.Conn, err error)
|
||||
// GetLastConnected returns last time the remote site was seen connected
|
||||
// GetLastConnected returns last time the cluster was seen connected
|
||||
GetLastConnected() time.Time
|
||||
// GetName returns site name (identified by authority domain's name)
|
||||
// GetName returns cluster name (identified by authority domain's name)
|
||||
GetName() string
|
||||
// GetStatus returns status of this site (either offline or connected)
|
||||
// GetStatus returns status of this cluster (either offline or connected)
|
||||
GetStatus() string
|
||||
// GetClient returns client connected to remote auth server
|
||||
GetClient() (authclient.ClientI, error)
|
||||
// CachingAccessPoint returns access point that is lightweight
|
||||
// but is resilient to auth server crashes
|
||||
CachingAccessPoint() (authclient.RemoteProxyAccessPoint, error)
|
||||
// NodeWatcher returns the node watcher that maintains the node set for the site
|
||||
// NodeWatcher returns the node watcher that maintains the node set for the cluster
|
||||
NodeWatcher() (*services.GenericWatcher[types.Server, readonly.Server], error)
|
||||
// GitServerWatcher returns the Git server watcher for the site
|
||||
// GitServerWatcher returns the Git server watcher for the cluster
|
||||
GitServerWatcher() (*services.GenericWatcher[types.Server, readonly.Server], error)
|
||||
// GetTunnelsCount returns the amount of active inbound tunnels
|
||||
// from the remote cluster
|
||||
GetTunnelsCount() int
|
||||
// IsClosed reports whether this RemoteSite has been closed and should no
|
||||
// IsClosed reports whether this Cluster has been closed and should no
|
||||
// longer be used.
|
||||
IsClosed() bool
|
||||
// Closer allows the site to be closed
|
||||
// Closer allows the Cluster to be closed
|
||||
io.Closer
|
||||
}
|
||||
|
||||
// Tunnel provides access to connected local or remote clusters
|
||||
// using unified interface.
|
||||
type Tunnel interface {
|
||||
// GetSites returns a list of connected remote sites
|
||||
GetSites() ([]RemoteSite, error)
|
||||
// GetSite returns remote site this node belongs to
|
||||
GetSite(domainName string) (RemoteSite, error)
|
||||
// GetSites returns a list of connected clusters
|
||||
GetSites() ([]Cluster, error)
|
||||
// GetSite returns cluster this node belongs to
|
||||
GetSite(domainName string) (Cluster, error)
|
||||
}
|
||||
|
||||
// Server is a TCP/IP SSH server which listens on an SSH endpoint and remote/local
|
||||
// sites connect and register with it.
|
||||
// cluster connect and register with it.
|
||||
type Server interface {
|
||||
Tunnel
|
||||
// Start starts server
|
||||
|
||||
@@ -57,13 +57,13 @@ type TunnelWithRoles struct {
|
||||
}
|
||||
|
||||
// GetSites returns a list of connected remote sites
|
||||
func (t *TunnelWithRoles) GetSites() ([]RemoteSite, error) {
|
||||
func (t *TunnelWithRoles) GetSites() ([]Cluster, error) {
|
||||
ctx := context.TODO()
|
||||
clusters, err := t.tunnel.GetSites()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
out := make([]RemoteSite, 0, len(clusters))
|
||||
out := make([]Cluster, 0, len(clusters))
|
||||
for _, cluster := range clusters {
|
||||
if t.localCluster == cluster.GetName() {
|
||||
out = append(out, cluster)
|
||||
@@ -89,7 +89,7 @@ func (t *TunnelWithRoles) GetSites() ([]RemoteSite, error) {
|
||||
}
|
||||
|
||||
// GetSite returns remote site this node belongs to
|
||||
func (t *TunnelWithRoles) GetSite(clusterName string) (RemoteSite, error) {
|
||||
func (t *TunnelWithRoles) GetSite(clusterName string) (Cluster, error) {
|
||||
ctx := context.TODO()
|
||||
cluster, err := t.tunnel.GetSite(clusterName)
|
||||
if err != nil {
|
||||
|
||||
@@ -31,47 +31,47 @@ import (
|
||||
// FakeServer is a fake Server implementation used in tests.
|
||||
type FakeServer struct {
|
||||
Server
|
||||
// Sites is a list of sites registered via this fake reverse tunnel.
|
||||
Sites []RemoteSite
|
||||
// Clusters is a list of clusters registered via this fake reverse tunnel.
|
||||
Clusters []Cluster
|
||||
}
|
||||
|
||||
// GetSites returns all available remote sites.
|
||||
func (s *FakeServer) GetSites() ([]RemoteSite, error) {
|
||||
return s.Sites, nil
|
||||
// GetSites returns all available clusters.
|
||||
func (s *FakeServer) GetSites() ([]Cluster, error) {
|
||||
return s.Clusters, nil
|
||||
}
|
||||
|
||||
// GetSite returns the remote site by name.
|
||||
func (s *FakeServer) GetSite(name string) (RemoteSite, error) {
|
||||
for _, site := range s.Sites {
|
||||
if site.GetName() == name {
|
||||
return site, nil
|
||||
// GetSite returns the cluster by name.
|
||||
func (s *FakeServer) GetSite(name string) (Cluster, error) {
|
||||
for _, cluster := range s.Clusters {
|
||||
if cluster.GetName() == name {
|
||||
return cluster, nil
|
||||
}
|
||||
}
|
||||
return nil, trace.NotFound("site %q not found", name)
|
||||
return nil, trace.NotFound("cluster %q not found", name)
|
||||
}
|
||||
|
||||
// FakeRemoteSite is a fake reversetunnelclient.RemoteSite implementation used in tests.
|
||||
type FakeRemoteSite struct {
|
||||
RemoteSite
|
||||
// Name is the remote site name.
|
||||
// FakeCluster is a fake reversetunnelclient.FakeCluster implementation used in tests.
|
||||
type FakeCluster struct {
|
||||
Cluster
|
||||
// Name is the cluster name.
|
||||
Name string
|
||||
// AccessPoint is the auth server client.
|
||||
AccessPoint authclient.RemoteProxyAccessPoint
|
||||
// OfflineTunnels is a list of server IDs that will return connection error.
|
||||
OfflineTunnels map[string]struct{}
|
||||
// connCh receives the connection when dialing this site.
|
||||
// connCh receives the connection when dialing this cluster.
|
||||
connCh chan net.Conn
|
||||
// connCounter count how many connection requests the remote received.
|
||||
connCounter int64
|
||||
// closedMtx is a mutex that protects closed.
|
||||
closedMtx sync.Mutex
|
||||
// closed is set to true after the site is being closed.
|
||||
// closed is set to true after the cluster is being closed.
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewFakeRemoteSite is a FakeRemoteSite constructor.
|
||||
func NewFakeRemoteSite(clusterName string, accessPoint authclient.RemoteProxyAccessPoint) *FakeRemoteSite {
|
||||
return &FakeRemoteSite{
|
||||
// NewFakeCluster is a FakeCluster constructor.
|
||||
func NewFakeCluster(clusterName string, accessPoint authclient.RemoteProxyAccessPoint) *FakeCluster {
|
||||
return &FakeCluster{
|
||||
Name: clusterName,
|
||||
connCh: make(chan net.Conn),
|
||||
AccessPoint: accessPoint,
|
||||
@@ -79,22 +79,22 @@ func NewFakeRemoteSite(clusterName string, accessPoint authclient.RemoteProxyAcc
|
||||
}
|
||||
|
||||
// CachingAccessPoint returns caching auth server client.
|
||||
func (s *FakeRemoteSite) CachingAccessPoint() (authclient.RemoteProxyAccessPoint, error) {
|
||||
func (s *FakeCluster) CachingAccessPoint() (authclient.RemoteProxyAccessPoint, error) {
|
||||
return s.AccessPoint, nil
|
||||
}
|
||||
|
||||
// GetName returns the remote site name.
|
||||
func (s *FakeRemoteSite) GetName() string {
|
||||
// GetName returns the remote cluster name.
|
||||
func (s *FakeCluster) GetName() string {
|
||||
return s.Name
|
||||
}
|
||||
|
||||
// ProxyConn returns proxy connection channel with incoming connections.
|
||||
func (s *FakeRemoteSite) ProxyConn() <-chan net.Conn {
|
||||
func (s *FakeCluster) ProxyConn() <-chan net.Conn {
|
||||
return s.connCh
|
||||
}
|
||||
|
||||
// Dial returns the connection to the remote site.
|
||||
func (s *FakeRemoteSite) Dial(params DialParams) (net.Conn, error) {
|
||||
// Dial returns the connection to the remote cluster.
|
||||
func (s *FakeCluster) Dial(params DialParams) (net.Conn, error) {
|
||||
atomic.AddInt64(&s.connCounter, 1)
|
||||
|
||||
if _, ok := s.OfflineTunnels[params.ServerID]; ok {
|
||||
@@ -114,7 +114,7 @@ func (s *FakeRemoteSite) Dial(params DialParams) (net.Conn, error) {
|
||||
return writerConn, nil
|
||||
}
|
||||
|
||||
func (s *FakeRemoteSite) Close() error {
|
||||
func (s *FakeCluster) Close() error {
|
||||
s.closedMtx.Lock()
|
||||
defer s.closedMtx.Unlock()
|
||||
close(s.connCh)
|
||||
@@ -122,6 +122,6 @@ func (s *FakeRemoteSite) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *FakeRemoteSite) DialCount() int64 {
|
||||
func (s *FakeCluster) DialCount() int64 {
|
||||
return atomic.LoadInt64(&s.connCounter)
|
||||
}
|
||||
|
||||
@@ -39,7 +39,7 @@ import (
|
||||
)
|
||||
|
||||
type sitesGetter interface {
|
||||
GetSites() ([]reversetunnelclient.RemoteSite, error)
|
||||
GetSites() ([]reversetunnelclient.Cluster, error)
|
||||
}
|
||||
|
||||
// NewAuthProxyDialerService create new instance of AuthProxyDialerService.
|
||||
|
||||
+18
-18
@@ -1475,19 +1475,19 @@ func TestRedisNil(t *testing.T) {
|
||||
}
|
||||
|
||||
type testContext struct {
|
||||
hostID string
|
||||
clusterName string
|
||||
tlsServer *authtest.TLSServer
|
||||
authServer *auth.Server
|
||||
authClient *authclient.Client
|
||||
proxyServer *ProxyServer
|
||||
mux *multiplexer.Mux
|
||||
mysqlListener net.Listener
|
||||
webListener *multiplexer.WebListener
|
||||
fakeRemoteSite *reversetunnelclient.FakeRemoteSite
|
||||
server *Server
|
||||
emitter *eventstest.ChannelEmitter
|
||||
databaseCA types.CertAuthority
|
||||
hostID string
|
||||
clusterName string
|
||||
tlsServer *authtest.TLSServer
|
||||
authServer *auth.Server
|
||||
authClient *authclient.Client
|
||||
proxyServer *ProxyServer
|
||||
mux *multiplexer.Mux
|
||||
mysqlListener net.Listener
|
||||
webListener *multiplexer.WebListener
|
||||
fakeCluster *reversetunnelclient.FakeCluster
|
||||
server *Server
|
||||
emitter *eventstest.ChannelEmitter
|
||||
databaseCA types.CertAuthority
|
||||
// postgres is a collection of Postgres databases the test uses.
|
||||
postgres map[string]testPostgres
|
||||
// mysql is a collection of MySQL databases the test uses.
|
||||
@@ -1630,7 +1630,7 @@ func (c *testContext) startHandlingConnections() {
|
||||
// Start all proxy services.
|
||||
c.startProxy()
|
||||
// Start handling database client connections on the database server.
|
||||
for conn := range c.fakeRemoteSite.ProxyConn() {
|
||||
for conn := range c.fakeCluster.ProxyConn() {
|
||||
go c.server.HandleConnection(conn)
|
||||
}
|
||||
}
|
||||
@@ -2381,11 +2381,11 @@ func setupTestContext(ctx context.Context, t testing.TB, withDatabases ...withDa
|
||||
}
|
||||
|
||||
// Establish fake reversetunnel b/w database proxy and database service.
|
||||
testCtx.fakeRemoteSite = reversetunnelclient.NewFakeRemoteSite(testCtx.clusterName, proxyAuthClient)
|
||||
t.Cleanup(func() { require.NoError(t, testCtx.fakeRemoteSite.Close()) })
|
||||
testCtx.fakeCluster = reversetunnelclient.NewFakeCluster(testCtx.clusterName, proxyAuthClient)
|
||||
t.Cleanup(func() { require.NoError(t, testCtx.fakeCluster.Close()) })
|
||||
tunnel := &reversetunnelclient.FakeServer{
|
||||
Sites: []reversetunnelclient.RemoteSite{
|
||||
testCtx.fakeRemoteSite,
|
||||
Clusters: []reversetunnelclient.Cluster{
|
||||
testCtx.fakeCluster,
|
||||
},
|
||||
}
|
||||
// Empty config means no limit.
|
||||
|
||||
@@ -519,7 +519,7 @@ func setupPostgres(ctx context.Context, t *testing.T, cfg *setupTLSTestCfg) *tes
|
||||
})
|
||||
|
||||
go func() {
|
||||
for conn := range testCtx.fakeRemoteSite.ProxyConn() {
|
||||
for conn := range testCtx.fakeCluster.ProxyConn() {
|
||||
go server1.HandleConnection(conn)
|
||||
}
|
||||
}()
|
||||
@@ -564,7 +564,7 @@ func setupMySQL(ctx context.Context, t *testing.T, cfg *setupTLSTestCfg) *testCo
|
||||
})
|
||||
|
||||
go func() {
|
||||
for conn := range testCtx.fakeRemoteSite.ProxyConn() {
|
||||
for conn := range testCtx.fakeCluster.ProxyConn() {
|
||||
go server1.HandleConnection(conn)
|
||||
}
|
||||
}()
|
||||
@@ -614,7 +614,7 @@ func setupMongo(ctx context.Context, t *testing.T, cfg *setupTLSTestCfg) *testCo
|
||||
})
|
||||
|
||||
go func() {
|
||||
for conn := range testCtx.fakeRemoteSite.ProxyConn() {
|
||||
for conn := range testCtx.fakeCluster.ProxyConn() {
|
||||
go server1.HandleConnection(conn)
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -61,7 +61,7 @@ type ProxyContext struct {
|
||||
// Identity is the authorized client Identity.
|
||||
Identity tlsca.Identity
|
||||
// Cluster is the remote Cluster running the database server.
|
||||
Cluster reversetunnelclient.RemoteSite
|
||||
Cluster reversetunnelclient.Cluster
|
||||
// Servers is a list of database Servers that proxy the requested database.
|
||||
Servers []types.DatabaseServer
|
||||
// AuthContext is a context of authenticated user.
|
||||
|
||||
@@ -60,7 +60,7 @@ func TestHA(t *testing.T) {
|
||||
Databases: types.Databases{offlineDB},
|
||||
HostID: offlineHostID,
|
||||
})
|
||||
testCtx.fakeRemoteSite.OfflineTunnels = map[string]struct{}{
|
||||
testCtx.fakeCluster.OfflineTunnels = map[string]struct{}{
|
||||
fmt.Sprintf("%v.%v", offlineHostID, testCtx.clusterName): {},
|
||||
}
|
||||
|
||||
@@ -78,7 +78,7 @@ func TestHA(t *testing.T) {
|
||||
HostID: onlineHostID,
|
||||
})
|
||||
go func() {
|
||||
for conn := range testCtx.fakeRemoteSite.ProxyConn() {
|
||||
for conn := range testCtx.fakeCluster.ProxyConn() {
|
||||
go onlineServer.HandleConnection(conn)
|
||||
}
|
||||
}()
|
||||
|
||||
+56
-56
@@ -1260,12 +1260,12 @@ func (h *Handler) handleGetUserOrResetToken(w http.ResponseWriter, r *http.Reque
|
||||
// getUserContext returns user context
|
||||
//
|
||||
// GET /webapi/sites/:site/context
|
||||
func (h *Handler) getUserContext(w http.ResponseWriter, r *http.Request, p httprouter.Params, c *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) getUserContext(w http.ResponseWriter, r *http.Request, p httprouter.Params, c *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
cn, err := h.cfg.AccessPoint.GetClusterName(r.Context())
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
if cn.GetClusterName() != site.GetName() {
|
||||
if cn.GetClusterName() != cluster.GetName() {
|
||||
return nil, trace.BadParameter("endpoint only implemented for root cluster")
|
||||
}
|
||||
accessChecker, err := c.GetUserAccessChecker()
|
||||
@@ -1317,7 +1317,7 @@ func (h *Handler) getUserContext(w http.ResponseWriter, r *http.Request, p httpr
|
||||
|
||||
userContext.AllowedSearchAsRoles = accessChecker.GetAllowedSearchAsRoles()
|
||||
|
||||
userContext.Cluster, err = ui.GetClusterDetails(r.Context(), site)
|
||||
userContext.Cluster, err = ui.GetClusterDetails(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3145,14 +3145,14 @@ type getClusterInfoResponse struct {
|
||||
}
|
||||
|
||||
// getClusterInfo returns the information about the cluster in the :site param
|
||||
func (h *Handler) getClusterInfo(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) getClusterInfo(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
clusterDetails, err := ui.GetClusterDetails(ctx, site)
|
||||
clusterDetails, err := ui.GetClusterDetails(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3179,7 +3179,7 @@ type getSiteNamespacesResponse struct {
|
||||
// Successful response:
|
||||
//
|
||||
// {"namespaces": [{..namespace resource...}]}
|
||||
func (h *Handler) getSiteNamespaces(w http.ResponseWriter, r *http.Request, _ httprouter.Params, c *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) getSiteNamespaces(w http.ResponseWriter, r *http.Request, _ httprouter.Params, c *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
return getSiteNamespacesResponse{
|
||||
Namespaces: []types.Namespace{types.DefaultNamespace()},
|
||||
}, nil
|
||||
@@ -3306,8 +3306,8 @@ func (h *Handler) getUserGroupLookup(ctx context.Context, clt apiclient.GetResou
|
||||
|
||||
// clusterUnifiedResourcesGet returns a list of resources for a given cluster site. This includes all resources available to be displayed in the web ui
|
||||
// such as Nodes, Apps, Desktops, etc etc
|
||||
func (h *Handler) clusterUnifiedResourcesGet(w http.ResponseWriter, request *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(request.Context(), site)
|
||||
func (h *Handler) clusterUnifiedResourcesGet(w http.ResponseWriter, request *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(request.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3344,9 +3344,9 @@ func (h *Handler) clusterUnifiedResourcesGet(w http.ResponseWriter, request *htt
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
unifiedResources = append(unifiedResources, ui.MakeServer(site.GetName(), r, logins, enriched.RequiresRequest))
|
||||
unifiedResources = append(unifiedResources, ui.MakeServer(cluster.GetName(), r, logins, enriched.RequiresRequest))
|
||||
case types.KindGitServer:
|
||||
unifiedResources = append(unifiedResources, ui.MakeGitServer(site.GetName(), r, enriched.RequiresRequest))
|
||||
unifiedResources = append(unifiedResources, ui.MakeGitServer(cluster.GetName(), r, enriched.RequiresRequest))
|
||||
}
|
||||
case types.DatabaseServer:
|
||||
db := ui.MakeDatabaseFromDatabaseServer(r, accessChecker, h.cfg.DatabaseREPLRegistry, enriched.RequiresRequest)
|
||||
@@ -3369,7 +3369,7 @@ func (h *Handler) clusterUnifiedResourcesGet(w http.ResponseWriter, request *htt
|
||||
app := ui.MakeApp(r.GetApp(), ui.MakeAppsConfig{
|
||||
LocalClusterName: h.auth.clusterName,
|
||||
LocalProxyDNSName: proxyDNSName,
|
||||
AppClusterName: site.GetName(),
|
||||
AppClusterName: cluster.GetName(),
|
||||
AllowedAWSRolesLookup: allowedAWSRolesLookup,
|
||||
UserGroupLookup: getUserGroupLookup(),
|
||||
Logger: h.logger,
|
||||
@@ -3382,7 +3382,7 @@ func (h *Handler) clusterUnifiedResourcesGet(w http.ResponseWriter, request *htt
|
||||
app := ui.MakeAppTypeFromSAMLApp(r, ui.MakeAppsConfig{
|
||||
LocalClusterName: h.auth.clusterName,
|
||||
LocalProxyDNSName: h.proxyDNSName(),
|
||||
AppClusterName: site.GetName(),
|
||||
AppClusterName: cluster.GetName(),
|
||||
RequiresRequest: enriched.RequiresRequest,
|
||||
})
|
||||
unifiedResources = append(unifiedResources, app)
|
||||
@@ -3408,10 +3408,10 @@ func (h *Handler) clusterUnifiedResourcesGet(w http.ResponseWriter, request *htt
|
||||
}
|
||||
|
||||
// clusterNodesGet returns a list of nodes for a given cluster site.
|
||||
func (h *Handler) clusterNodesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) clusterNodesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
// Get a client to the Auth Server with the logged in user's identity. The
|
||||
// identity of the logged in user is used to fetch the list of nodes.
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3444,7 +3444,7 @@ func (h *Handler) clusterNodesGet(w http.ResponseWriter, r *http.Request, p http
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
uiServers = append(uiServers, ui.MakeServer(site.GetName(), server, logins, false /* requiresRequest */))
|
||||
uiServers = append(uiServers, ui.MakeServer(cluster.GetName(), server, logins, false /* requiresRequest */))
|
||||
}
|
||||
|
||||
return listResourcesGetResponse{
|
||||
@@ -3458,8 +3458,8 @@ func (h *Handler) clusterNodesGet(w http.ResponseWriter, r *http.Request, p http
|
||||
const iso8601MilliFormat = "2006-01-02T15:04:05.000Z0700"
|
||||
|
||||
// notificationsGet returns a paginated list of notifications for a user.
|
||||
func (h *Handler) notificationsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) notificationsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3501,8 +3501,8 @@ type GetNotificationsResponse struct {
|
||||
}
|
||||
|
||||
// notificationsUpsertLastSeenTimestamp upserts a user's last seen notification timestamp.
|
||||
func (h *Handler) notificationsUpsertLastSeenTimestamp(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) notificationsUpsertLastSeenTimestamp(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3538,8 +3538,8 @@ type UpsertUserLastSeenNotificationRequest struct {
|
||||
}
|
||||
|
||||
// notificationsUpsertNotificationState upserts a user notification state.
|
||||
func (h *Handler) notificationsUpsertNotificationState(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) notificationsUpsertNotificationState(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3580,10 +3580,10 @@ type getLoginAlertsResponse struct {
|
||||
}
|
||||
|
||||
// clusterLoginAlertsGet returns a list of on-login alerts for the user.
|
||||
func (h *Handler) clusterLoginAlertsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) clusterLoginAlertsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
// Get a client to the Auth Server with the logged in user's identity. The
|
||||
// identity of the logged in user is used to fetch the list of alerts.
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3607,10 +3607,10 @@ func (h *Handler) getClusterLocks(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sessionCtx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
) (any, error) {
|
||||
ctx := r.Context()
|
||||
clt, err := sessionCtx.GetUserClient(ctx, site)
|
||||
clt, err := sessionCtx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3634,7 +3634,7 @@ func (h *Handler) createClusterLock(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sessionCtx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
) (any, error) {
|
||||
var req *createLockReq
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
@@ -3642,7 +3642,7 @@ func (h *Handler) createClusterLock(
|
||||
}
|
||||
|
||||
ctx := r.Context()
|
||||
clt, err := sessionCtx.GetUserClient(ctx, site)
|
||||
clt, err := sessionCtx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3682,10 +3682,10 @@ func (h *Handler) deleteClusterLock(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sessionCtx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
) (any, error) {
|
||||
ctx := r.Context()
|
||||
clt, err := sessionCtx.GetUserClient(ctx, site)
|
||||
clt, err := sessionCtx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3735,7 +3735,7 @@ func (h *Handler) siteNodeConnect(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sessionCtx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
ws *websocket.Conn,
|
||||
) (any, error) {
|
||||
q := r.URL.Query()
|
||||
@@ -3748,7 +3748,7 @@ func (h *Handler) siteNodeConnect(
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sessionCtx.GetUserClient(r.Context(), site)
|
||||
clt, err := sessionCtx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3764,7 +3764,7 @@ func (h *Handler) siteNodeConnect(
|
||||
tracker types.SessionTracker
|
||||
)
|
||||
|
||||
clusterName := site.GetName()
|
||||
clusterName := cluster.GetName()
|
||||
if req.SessionID.IsZero() {
|
||||
// An existing session ID was not provided so we need to create a new one.
|
||||
sessionData, err = h.generateSession(r.Context(), &req, clusterName, sessionCtx)
|
||||
@@ -3797,7 +3797,7 @@ func (h *Handler) siteNodeConnect(
|
||||
"websid", sessionCtx.GetSessionID(),
|
||||
)
|
||||
|
||||
authAccessPoint, err := site.CachingAccessPoint()
|
||||
authAccessPoint, err := cluster.CachingAccessPoint()
|
||||
if err != nil {
|
||||
h.logger.DebugContext(r.Context(), "Unable to get auth access point", "error", err)
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -3818,7 +3818,7 @@ func (h *Handler) siteNodeConnect(
|
||||
keepAliveInterval = req.KeepAliveInterval
|
||||
}
|
||||
|
||||
nw, err := site.NodeWatcher()
|
||||
nw, err := cluster.NodeWatcher()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -3912,7 +3912,7 @@ func (h *Handler) podConnect(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
ws *websocket.Conn,
|
||||
) (any, error) {
|
||||
q := r.URL.Query()
|
||||
@@ -3932,7 +3932,7 @@ func (h *Handler) podConnect(
|
||||
params.SessionID.String(),
|
||||
params.ParticipantMode,
|
||||
sctx,
|
||||
site,
|
||||
cluster,
|
||||
ws,
|
||||
))
|
||||
}
|
||||
@@ -3959,7 +3959,7 @@ func (h *Handler) podConnect(
|
||||
sess := session.Session{
|
||||
Kind: types.KubernetesSessionKind,
|
||||
Login: "root",
|
||||
ClusterName: site.GetName(),
|
||||
ClusterName: cluster.GetName(),
|
||||
KubernetesClusterName: execReq.KubeCluster,
|
||||
ID: session.NewID(),
|
||||
Created: h.clock.Now().UTC(),
|
||||
@@ -3977,7 +3977,7 @@ func (h *Handler) podConnect(
|
||||
"websid", sctx.GetSessionID(),
|
||||
)
|
||||
|
||||
authAccessPoint, err := site.CachingAccessPoint()
|
||||
authAccessPoint, err := cluster.CachingAccessPoint()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -4000,7 +4000,7 @@ func (h *Handler) podConnect(
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -4009,7 +4009,7 @@ func (h *Handler) podConnect(
|
||||
req: execReq,
|
||||
sess: sess,
|
||||
sctx: sctx,
|
||||
teleportCluster: site.GetName(),
|
||||
teleportCluster: cluster.GetName(),
|
||||
ws: ws,
|
||||
keepAliveInterval: netConfig.GetKeepAliveInterval(),
|
||||
logger: h.logger.With(teleport.ComponentKey, "pod"),
|
||||
@@ -4192,8 +4192,8 @@ func trackerToLegacySession(tracker types.SessionTracker, clusterName string) se
|
||||
// clusterActiveAndPendingSessionsGet gets the list of active and pending sessions for a site.
|
||||
//
|
||||
// GET /v1/webapi/sites/:site/sessions
|
||||
func (h *Handler) clusterActiveAndPendingSessionsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) clusterActiveAndPendingSessionsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -4256,7 +4256,7 @@ func toFieldsSlice(rawEvents []apievents.AuditEvent) ([]events.EventFields, erro
|
||||
// "order": optional ordering of events. Can be either "asc" or "desc"
|
||||
// for ascending and descending respectively.
|
||||
// If no order is provided it defaults to descending.
|
||||
func (h *Handler) clusterSearchEvents(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) clusterSearchEvents(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
values := r.URL.Query()
|
||||
|
||||
var eventTypes []string
|
||||
@@ -4274,7 +4274,7 @@ func (h *Handler) clusterSearchEvents(w http.ResponseWriter, r *http.Request, p
|
||||
StartKey: startKey,
|
||||
})
|
||||
}
|
||||
return clusterEventsList(r.Context(), sctx, site, r.URL.Query(), searchEvents)
|
||||
return clusterEventsList(r.Context(), sctx, cluster, r.URL.Query(), searchEvents)
|
||||
}
|
||||
|
||||
// clusterSearchSessionEvents returns session events matching the criteria.
|
||||
@@ -4291,7 +4291,7 @@ func (h *Handler) clusterSearchEvents(w http.ResponseWriter, r *http.Request, p
|
||||
// "order": optional ordering of events. Can be either "asc" or "desc"
|
||||
// for ascending and descending respectively.
|
||||
// If no order is provided it defaults to descending.
|
||||
func (h *Handler) clusterSearchSessionEvents(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) clusterSearchSessionEvents(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
searchSessionEvents := func(clt authclient.ClientI, from, to time.Time, limit int, order types.EventOrder, startKey string) ([]apievents.AuditEvent, string, error) {
|
||||
return clt.SearchSessionEvents(r.Context(), events.SearchSessionEventsRequest{
|
||||
From: from,
|
||||
@@ -4301,12 +4301,12 @@ func (h *Handler) clusterSearchSessionEvents(w http.ResponseWriter, r *http.Requ
|
||||
StartKey: startKey,
|
||||
})
|
||||
}
|
||||
return clusterEventsList(r.Context(), sctx, site, r.URL.Query(), searchSessionEvents)
|
||||
return clusterEventsList(r.Context(), sctx, cluster, r.URL.Query(), searchSessionEvents)
|
||||
}
|
||||
|
||||
// clusterEventsList returns a list of audit events obtained using the provided
|
||||
// searchEvents method.
|
||||
func clusterEventsList(ctx context.Context, sctx *SessionContext, site reversetunnelclient.RemoteSite, values url.Values, searchEvents func(clt authclient.ClientI, from, to time.Time, limit int, order types.EventOrder, startKey string) ([]apievents.AuditEvent, string, error)) (any, error) {
|
||||
func clusterEventsList(ctx context.Context, sctx *SessionContext, cluster reversetunnelclient.Cluster, values url.Values, searchEvents func(clt authclient.ClientI, from, to time.Time, limit int, order types.EventOrder, startKey string) ([]apievents.AuditEvent, string, error)) (any, error) {
|
||||
from, err := queryTime(values, "from", time.Now().UTC().AddDate(0, -1, 0))
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -4329,7 +4329,7 @@ func clusterEventsList(ctx context.Context, sctx *SessionContext, site reversetu
|
||||
|
||||
startKey := values.Get("startKey")
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -4560,10 +4560,10 @@ const currentSiteShortcut = "-current-"
|
||||
type ContextHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext) (any, error)
|
||||
|
||||
// ClusterHandler is a authenticated handler that is called for some existing remote cluster
|
||||
type ClusterHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error)
|
||||
type ClusterHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error)
|
||||
|
||||
// ClusterWebsocketHandler is a authenticated websocket handler that is called for some existing remote cluster
|
||||
type ClusterWebsocketHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite, ws *websocket.Conn) (any, error)
|
||||
type ClusterWebsocketHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster, ws *websocket.Conn) (any, error)
|
||||
|
||||
// WithClusterAuth wraps a ClusterHandler to ensure that a request is authenticated to this proxy
|
||||
// (the same as WithAuth), as well as to grab the remoteSite (which can represent this local cluster
|
||||
@@ -4641,7 +4641,7 @@ func (h *Handler) WithClusterAuthWebSocket(fn ClusterWebsocketHandler) httproute
|
||||
// *SessionContext (same as AuthenticateRequest), and also grabs the
|
||||
// remoteSite (which can represent this local cluster or a remote
|
||||
// trusted cluster) as specified by the ":site" url parameter.
|
||||
func (h *Handler) authenticateWSRequestWithCluster(w http.ResponseWriter, r *http.Request, p httprouter.Params) (*SessionContext, *websocket.Conn, reversetunnelclient.RemoteSite, error) {
|
||||
func (h *Handler) authenticateWSRequestWithCluster(w http.ResponseWriter, r *http.Request, p httprouter.Params) (*SessionContext, *websocket.Conn, reversetunnelclient.Cluster, error) {
|
||||
sctx, ws, err := h.AuthenticateRequestWS(w, r)
|
||||
if err != nil {
|
||||
return nil, nil, nil, trace.Wrap(err)
|
||||
@@ -4659,7 +4659,7 @@ func (h *Handler) authenticateWSRequestWithCluster(w http.ResponseWriter, r *htt
|
||||
// to this proxy, returning the *SessionContext (same as AuthenticateRequest),
|
||||
// and also grabs the remoteSite (which can represent this local cluster or a
|
||||
// remote trusted cluster) as specified by the ":site" url parameter.
|
||||
func (h *Handler) authenticateRequestWithCluster(w http.ResponseWriter, r *http.Request, p httprouter.Params) (*SessionContext, reversetunnelclient.RemoteSite, error) {
|
||||
func (h *Handler) authenticateRequestWithCluster(w http.ResponseWriter, r *http.Request, p httprouter.Params) (*SessionContext, reversetunnelclient.Cluster, error) {
|
||||
sctx, err := h.AuthenticateRequest(w, r, true)
|
||||
if err != nil {
|
||||
return nil, nil, trace.Wrap(err)
|
||||
@@ -4675,7 +4675,7 @@ func (h *Handler) authenticateRequestWithCluster(w http.ResponseWriter, r *http.
|
||||
|
||||
// getSiteByParams gets the remoteSite (which can represent this local cluster or a
|
||||
// remote trusted cluster) as specified by the ":site" url parameter.
|
||||
func (h *Handler) getSiteByParams(ctx context.Context, sctx *SessionContext, p httprouter.Params) (reversetunnelclient.RemoteSite, error) {
|
||||
func (h *Handler) getSiteByParams(ctx context.Context, sctx *SessionContext, p httprouter.Params) (reversetunnelclient.Cluster, error) {
|
||||
clusterName := p.ByName("site")
|
||||
site, err := h.getSiteByClusterName(ctx, sctx, clusterName)
|
||||
if err != nil {
|
||||
@@ -4685,7 +4685,7 @@ func (h *Handler) getSiteByParams(ctx context.Context, sctx *SessionContext, p h
|
||||
return site, nil
|
||||
}
|
||||
|
||||
func (h *Handler) getSiteByClusterName(ctx context.Context, sctx *SessionContext, clusterName string) (reversetunnelclient.RemoteSite, error) {
|
||||
func (h *Handler) getSiteByClusterName(ctx context.Context, sctx *SessionContext, clusterName string) (reversetunnelclient.Cluster, error) {
|
||||
if clusterName == currentSiteShortcut {
|
||||
res, err := h.cfg.ProxyClient.GetClusterName(ctx)
|
||||
if err != nil {
|
||||
@@ -4760,7 +4760,7 @@ func (h *Handler) WithClusterClientProvider(fn ClusterClientHandler) httprouter.
|
||||
}
|
||||
|
||||
// ProvisionTokenHandler is a authenticated handler that is called for some existing Token
|
||||
type ProvisionTokenHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, site reversetunnelclient.RemoteSite, token types.ProvisionToken) (any, error)
|
||||
type ProvisionTokenHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, cluster reversetunnelclient.Cluster, token types.ProvisionToken) (any, error)
|
||||
|
||||
// WithProvisionTokenAuth ensures that request is authenticated with a provision token.
|
||||
// Provision tokens, when used like this are invalidated as soon as used.
|
||||
|
||||
+32
-32
@@ -327,10 +327,10 @@ func TestMatchApplicationServers(t *testing.T) {
|
||||
// Create a httptest server to serve the application requests. It must serve
|
||||
// TLS content with the generated certificate.
|
||||
expectedContent := "Hello application"
|
||||
fakeRemoteSite := startFakeAppServerOnRemoteSite(t, clusterName, authClient, cert, key)
|
||||
fakeCluster := startFakeAppServerOnCluster(t, clusterName, authClient, cert, key)
|
||||
tunnel := &reversetunnelclient.FakeServer{
|
||||
Sites: []reversetunnelclient.RemoteSite{
|
||||
fakeRemoteSite,
|
||||
Clusters: []reversetunnelclient.Cluster{
|
||||
fakeCluster,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -347,9 +347,9 @@ func TestMatchApplicationServers(t *testing.T) {
|
||||
})
|
||||
|
||||
require.Equal(t, http.StatusOK, status)
|
||||
// Remote site should receive only 4 connection requests: 3 from the
|
||||
// Cluster should receive only 4 connection requests: 3 from the
|
||||
// MatchHealthy and 1 from the transport.
|
||||
require.Equal(t, int64(4), fakeRemoteSite.DialCount())
|
||||
require.Equal(t, int64(4), fakeCluster.DialCount())
|
||||
// Guarantee the request was returned by the httptest server.
|
||||
require.Equal(t, expectedContent, content)
|
||||
}
|
||||
@@ -369,14 +369,14 @@ func TestHealthCheckAppServer(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
desc string
|
||||
publicAddr string
|
||||
appServersFunc func(t *testing.T, remoteSite *reversetunnelclient.FakeRemoteSite) []types.AppServer
|
||||
appServersFunc func(t *testing.T, cluster *reversetunnelclient.FakeCluster) []types.AppServer
|
||||
expectedTunnelCalls int
|
||||
expectErr require.ErrorAssertionFunc
|
||||
}{
|
||||
{
|
||||
desc: "match and online services",
|
||||
publicAddr: "valid.example.com",
|
||||
appServersFunc: func(t *testing.T, _ *reversetunnelclient.FakeRemoteSite) []types.AppServer {
|
||||
appServersFunc: func(t *testing.T, _ *reversetunnelclient.FakeCluster) []types.AppServer {
|
||||
return []types.AppServer{createAppServer(t, "valid.example.com")}
|
||||
},
|
||||
expectedTunnelCalls: 1,
|
||||
@@ -385,9 +385,9 @@ func TestHealthCheckAppServer(t *testing.T) {
|
||||
{
|
||||
desc: "match and but no online services",
|
||||
publicAddr: "valid.example.com",
|
||||
appServersFunc: func(t *testing.T, tunnel *reversetunnelclient.FakeRemoteSite) []types.AppServer {
|
||||
appServersFunc: func(t *testing.T, cluster *reversetunnelclient.FakeCluster) []types.AppServer {
|
||||
appServer := createAppServer(t, "valid.example.com")
|
||||
tunnel.OfflineTunnels = map[string]struct{}{
|
||||
cluster.OfflineTunnels = map[string]struct{}{
|
||||
fmt.Sprintf("%s.%s", appServer.GetHostID(), clusterName): {},
|
||||
}
|
||||
return []types.AppServer{appServer}
|
||||
@@ -398,7 +398,7 @@ func TestHealthCheckAppServer(t *testing.T) {
|
||||
{
|
||||
desc: "no match",
|
||||
publicAddr: "valid.example.com",
|
||||
appServersFunc: func(t *testing.T, tunnel *reversetunnelclient.FakeRemoteSite) []types.AppServer {
|
||||
appServersFunc: func(t *testing.T, _ *reversetunnelclient.FakeCluster) []types.AppServer {
|
||||
return []types.AppServer{}
|
||||
},
|
||||
expectedTunnelCalls: 0,
|
||||
@@ -415,11 +415,11 @@ func TestHealthCheckAppServer(t *testing.T) {
|
||||
caCert: cert,
|
||||
}
|
||||
|
||||
fakeRemoteSite := startFakeAppServerOnRemoteSite(t, clusterName, authClient, cert, key)
|
||||
authClient.appServers = tc.appServersFunc(t, fakeRemoteSite)
|
||||
fakeCluster := startFakeAppServerOnCluster(t, clusterName, authClient, cert, key)
|
||||
authClient.appServers = tc.appServersFunc(t, fakeCluster)
|
||||
|
||||
tunnel := &reversetunnelclient.FakeServer{
|
||||
Sites: []reversetunnelclient.RemoteSite{fakeRemoteSite},
|
||||
Clusters: []reversetunnelclient.Cluster{fakeCluster},
|
||||
}
|
||||
|
||||
appHandler, err := NewHandler(ctx, &HandlerConfig{
|
||||
@@ -434,7 +434,7 @@ func TestHealthCheckAppServer(t *testing.T) {
|
||||
|
||||
err = appHandler.HealthCheckAppServer(ctx, tc.publicAddr, clusterName)
|
||||
tc.expectErr(t, err)
|
||||
require.Equal(t, int64(tc.expectedTunnelCalls), fakeRemoteSite.DialCount())
|
||||
require.Equal(t, int64(tc.expectedTunnelCalls), fakeCluster.DialCount())
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -576,26 +576,26 @@ func (c *mockAuthClient) GetProxies() ([]types.Server, error) {
|
||||
return []types.Server{}, nil
|
||||
}
|
||||
|
||||
// fakeRemoteListener Implements a `net.Listener` that return `net.Conn` from
|
||||
// the `FakeRemoteSite`.
|
||||
type fakeRemoteListener struct {
|
||||
fakeRemote *reversetunnelclient.FakeRemoteSite
|
||||
// fakeClusterListener Implements a `net.Listener` that return `net.Conn` from
|
||||
// the `FakeCluster`.
|
||||
type fakeClusterListener struct {
|
||||
fakeCluster *reversetunnelclient.FakeCluster
|
||||
}
|
||||
|
||||
func (r *fakeRemoteListener) Accept() (net.Conn, error) {
|
||||
conn, ok := <-r.fakeRemote.ProxyConn()
|
||||
func (r *fakeClusterListener) Accept() (net.Conn, error) {
|
||||
conn, ok := <-r.fakeCluster.ProxyConn()
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("remote closed")
|
||||
return nil, fmt.Errorf("cluster closed")
|
||||
}
|
||||
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (r *fakeRemoteListener) Close() error {
|
||||
func (r *fakeClusterListener) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *fakeRemoteListener) Addr() net.Addr {
|
||||
func (r *fakeClusterListener) Addr() net.Addr {
|
||||
return &net.IPAddr{}
|
||||
}
|
||||
|
||||
@@ -774,19 +774,19 @@ func TestMakeAppRedirectURL(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func startFakeAppServerOnRemoteSite(t *testing.T, clusterName string, accessPoint authclient.RemoteProxyAccessPoint, cert, key []byte) *reversetunnelclient.FakeRemoteSite {
|
||||
func startFakeAppServerOnCluster(t *testing.T, clusterName string, accessPoint authclient.RemoteProxyAccessPoint, cert, key []byte) *reversetunnelclient.FakeCluster {
|
||||
t.Helper()
|
||||
|
||||
tlsCert, err := tls.X509KeyPair(cert, key)
|
||||
require.NoError(t, err)
|
||||
|
||||
fakeRemoteSite := reversetunnelclient.NewFakeRemoteSite(clusterName, accessPoint)
|
||||
fakeCluster := reversetunnelclient.NewFakeCluster(clusterName, accessPoint)
|
||||
server := &httptest.Server{
|
||||
TLS: &tls.Config{
|
||||
Certificates: []tls.Certificate{tlsCert},
|
||||
},
|
||||
Listener: &fakeRemoteListener{
|
||||
fakeRemote: fakeRemoteSite,
|
||||
Listener: &fakeClusterListener{
|
||||
fakeCluster: fakeCluster,
|
||||
},
|
||||
Config: &http.Server{Handler: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
fmt.Fprint(w, "Hello application")
|
||||
@@ -794,11 +794,11 @@ func startFakeAppServerOnRemoteSite(t *testing.T, clusterName string, accessPoin
|
||||
}
|
||||
server.StartTLS()
|
||||
t.Cleanup(func() {
|
||||
// Close fake remote site first to make sure fake listener quits.
|
||||
fakeRemoteSite.Close()
|
||||
// Close fake cluster first to make sure fake listener quits.
|
||||
fakeCluster.Close()
|
||||
server.Close()
|
||||
})
|
||||
return fakeRemoteSite
|
||||
return fakeCluster
|
||||
}
|
||||
|
||||
func TestHandlerAuthenticate(t *testing.T) {
|
||||
@@ -825,14 +825,14 @@ func TestHandlerAuthenticate(t *testing.T) {
|
||||
caCert: cert,
|
||||
}
|
||||
|
||||
fakeRemoteSite := startFakeAppServerOnRemoteSite(t, clusterName, authClient, cert, key)
|
||||
fakeCluster := startFakeAppServerOnCluster(t, clusterName, authClient, cert, key)
|
||||
|
||||
appHandler, err := NewHandler(ctx, &HandlerConfig{
|
||||
Clock: fakeClock,
|
||||
AuthClient: authClient,
|
||||
AccessPoint: authClient,
|
||||
ProxyClient: &reversetunnelclient.FakeServer{
|
||||
Sites: []reversetunnelclient.RemoteSite{fakeRemoteSite},
|
||||
Clusters: []reversetunnelclient.Cluster{fakeCluster},
|
||||
},
|
||||
CipherSuites: utils.DefaultCipherSuites(),
|
||||
IntegrationAppHandler: &mockIntegrationAppHandler{},
|
||||
|
||||
@@ -77,7 +77,7 @@ func TestMatchHealthy(t *testing.T) {
|
||||
for name, test := range testCases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
match := MatchHealthy(&mockProxyClient{
|
||||
remoteSite: &mockRemoteSite{
|
||||
cluster: &mockCluster{
|
||||
dialErr: test.dialErr,
|
||||
},
|
||||
}, "")
|
||||
@@ -113,19 +113,19 @@ func mustNewAppServer(t *testing.T, origin string) func() types.AppServer {
|
||||
|
||||
type mockProxyClient struct {
|
||||
reversetunnelclient.Tunnel
|
||||
remoteSite *mockRemoteSite
|
||||
cluster *mockCluster
|
||||
}
|
||||
|
||||
func (p *mockProxyClient) GetSite(_ string) (reversetunnelclient.RemoteSite, error) {
|
||||
return p.remoteSite, nil
|
||||
func (p *mockProxyClient) GetSite(_ string) (reversetunnelclient.Cluster, error) {
|
||||
return p.cluster, nil
|
||||
}
|
||||
|
||||
type mockRemoteSite struct {
|
||||
reversetunnelclient.RemoteSite
|
||||
type mockCluster struct {
|
||||
reversetunnelclient.Cluster
|
||||
dialErr error
|
||||
}
|
||||
|
||||
func (r *mockRemoteSite) Dial(_ reversetunnelclient.DialParams) (net.Conn, error) {
|
||||
func (r *mockCluster) Dial(_ reversetunnelclient.DialParams) (net.Conn, error) {
|
||||
if r.dialErr != nil {
|
||||
return nil, r.dialErr
|
||||
}
|
||||
|
||||
@@ -220,12 +220,12 @@ func Test_transport_rewriteRedirect(t *testing.T) {
|
||||
type fakeTunnel struct {
|
||||
reversetunnelclient.Tunnel
|
||||
|
||||
fakeSite *reversetunnelclient.FakeRemoteSite
|
||||
err error
|
||||
fakeCluster *reversetunnelclient.FakeCluster
|
||||
err error
|
||||
}
|
||||
|
||||
func (f fakeTunnel) GetSite(domainName string) (reversetunnelclient.RemoteSite, error) {
|
||||
return f.fakeSite, f.err
|
||||
func (f fakeTunnel) GetSite(domainName string) (reversetunnelclient.Cluster, error) {
|
||||
return f.fakeCluster, f.err
|
||||
}
|
||||
|
||||
func TestTransport_DialContextNoServersAvailable(t *testing.T) {
|
||||
|
||||
@@ -31,8 +31,8 @@ import (
|
||||
)
|
||||
|
||||
// getConnectionDiagnostic returns a connection diagnostic connection diagnostics.
|
||||
func (h *Handler) getConnectionDiagnostic(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) getConnectionDiagnostic(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -52,7 +52,7 @@ func (h *Handler) getConnectionDiagnostic(w http.ResponseWriter, r *http.Request
|
||||
}
|
||||
|
||||
// diagnoseConnection executes and returns a connection diagnostic.
|
||||
func (h *Handler) diagnoseConnection(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) diagnoseConnection(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
req := conntest.TestConnectionRequest{}
|
||||
if err := httplib.ReadJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -62,7 +62,7 @@ func (h *Handler) diagnoseConnection(w http.ResponseWriter, r *http.Request, p h
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
userClt, err := sctx.GetUserClient(r.Context(), site)
|
||||
userClt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+10
-10
@@ -118,7 +118,7 @@ func (r *createOrOverwriteDatabaseRequest) checkAndSetDefaults() error {
|
||||
}
|
||||
|
||||
// handleDatabaseCreate creates a database's metadata.
|
||||
func (h *Handler) handleDatabaseCreateOrOverwrite(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) handleDatabaseCreateOrOverwrite(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var req *createOrOverwriteDatabaseRequest
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -133,7 +133,7 @@ func (h *Handler) handleDatabaseCreateOrOverwrite(w http.ResponseWriter, r *http
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -199,7 +199,7 @@ func (r *updateDatabaseRequest) checkAndSetDefaults() error {
|
||||
}
|
||||
|
||||
// handleDatabaseUpdate updates the database
|
||||
func (h *Handler) handleDatabasePartialUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) handleDatabasePartialUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
databaseName := p.ByName("database")
|
||||
if databaseName == "" {
|
||||
return nil, trace.BadParameter("a database name is required")
|
||||
@@ -214,7 +214,7 @@ func (h *Handler) handleDatabasePartialUpdate(w http.ResponseWriter, r *http.Req
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -292,13 +292,13 @@ type databaseIAMPolicyAWS struct {
|
||||
}
|
||||
|
||||
// handleDatabaseGetIAMPolicy returns the required IAM policy for database.
|
||||
func (h *Handler) handleDatabaseGetIAMPolicy(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) handleDatabaseGetIAMPolicy(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
databaseName := p.ByName("database")
|
||||
if databaseName == "" {
|
||||
return nil, trace.BadParameter("missing database name")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -411,7 +411,7 @@ func (h *Handler) dbConnect(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
ws *websocket.Conn,
|
||||
) (any, error) {
|
||||
// Create a context for signaling when the terminal session is over and
|
||||
@@ -469,7 +469,7 @@ func (h *Handler) dbConnect(
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -479,7 +479,7 @@ func (h *Handler) dbConnect(
|
||||
req: req,
|
||||
ws: ws,
|
||||
sctx: sctx,
|
||||
site: site,
|
||||
site: cluster,
|
||||
clt: clt,
|
||||
keepAliveInterval: netConfig.GetKeepAliveInterval(),
|
||||
registry: h.cfg.DatabaseREPLRegistry,
|
||||
@@ -563,7 +563,7 @@ type databaseInteractiveSessionConfig struct {
|
||||
log *slog.Logger
|
||||
req *DatabaseSessionRequest
|
||||
sctx *SessionContext
|
||||
site reversetunnelclient.RemoteSite
|
||||
site reversetunnelclient.Cluster
|
||||
clt authclient.ClientI
|
||||
keepAliveInterval time.Duration
|
||||
registry dbrepl.REPLRegistry
|
||||
|
||||
+10
-10
@@ -57,7 +57,7 @@ func (h *Handler) desktopConnectHandle(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
ws *websocket.Conn,
|
||||
) (any, error) {
|
||||
desktopName := p.ByName("desktopName")
|
||||
@@ -67,11 +67,11 @@ func (h *Handler) desktopConnectHandle(
|
||||
|
||||
log := sctx.cfg.Log.With(
|
||||
"desktop_name", desktopName,
|
||||
"cluster_name", site.GetName(),
|
||||
"cluster_name", cluster.GetName(),
|
||||
)
|
||||
log.DebugContext(r.Context(), "New desktop access websocket connection")
|
||||
|
||||
if err := h.createDesktopConnection(r, desktopName, site.GetName(), log, sctx, site, ws); err != nil {
|
||||
if err := h.createDesktopConnection(r, desktopName, cluster.GetName(), log, sctx, cluster, ws); err != nil {
|
||||
// createDesktopConnection makes a best effort attempt to send an error to the user
|
||||
// (via websocket) before terminating the connection. We log the error here, but
|
||||
// return nil because our HTTP middleware will try to write the returned error in JSON
|
||||
@@ -88,7 +88,7 @@ func (h *Handler) createDesktopConnection(
|
||||
clusterName string,
|
||||
log *slog.Logger,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
ws *websocket.Conn,
|
||||
) error {
|
||||
defer ws.Close()
|
||||
@@ -155,7 +155,7 @@ func (h *Handler) createDesktopConnection(
|
||||
withheld = append(withheld, msg)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return sendTDPError(trace.Wrap(err))
|
||||
}
|
||||
@@ -167,7 +167,7 @@ func (h *Handler) createDesktopConnection(
|
||||
}
|
||||
|
||||
// Check if MFA is required and create a UserCertsRequest.
|
||||
mfaRequired, certsReq, err := h.prepareForCertIssuance(ctx, sctx, site, pk.Public(), desktopName, username)
|
||||
mfaRequired, certsReq, err := h.prepareForCertIssuance(ctx, sctx, cluster, pk.Public(), desktopName, username)
|
||||
if err != nil {
|
||||
return sendTDPError(err)
|
||||
}
|
||||
@@ -189,7 +189,7 @@ func (h *Handler) createDesktopConnection(
|
||||
serviceConn, version, err := desktop.ConnectToWindowsService(ctx, &desktop.ConnectionConfig{
|
||||
Log: log,
|
||||
DesktopsGetter: clt,
|
||||
Site: site,
|
||||
Site: cluster,
|
||||
ClientSrcAddr: clientSrcAddr,
|
||||
ClientDstAddr: clientDstAddr,
|
||||
DesktopName: desktopName,
|
||||
@@ -297,7 +297,7 @@ func createUserCertsRequest(
|
||||
func (h *Handler) prepareForCertIssuance(
|
||||
ctx context.Context,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
publicKey crypto.PublicKey,
|
||||
desktopName, username string,
|
||||
) (mfaRequired bool, certsReq *proto.UserCertsRequest, err error) {
|
||||
@@ -307,12 +307,12 @@ func (h *Handler) prepareForCertIssuance(
|
||||
DesktopName: desktopName,
|
||||
Login: username,
|
||||
},
|
||||
}, sctx, site)
|
||||
}, sctx, cluster)
|
||||
if err != nil {
|
||||
return false, nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
certsReq, err = createUserCertsRequest(sctx, publicKey, desktopName, username, site.GetName())
|
||||
certsReq, err = createUserCertsRequest(sctx, publicKey, desktopName, username, cluster.GetName())
|
||||
if err != nil {
|
||||
return false, nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -37,7 +37,7 @@ func (h *Handler) desktopPlaybackHandle(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
ws *websocket.Conn,
|
||||
) (any, error) {
|
||||
sID := p.ByName("sid")
|
||||
@@ -45,7 +45,7 @@ func (h *Handler) desktopPlaybackHandle(
|
||||
return nil, trace.BadParameter("missing session ID in request URL")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+10
-10
@@ -33,7 +33,7 @@ import (
|
||||
)
|
||||
|
||||
// discoveryconfigCreate creates a DiscoveryConfig
|
||||
func (h *Handler) discoveryconfigCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) discoveryconfigCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var req ui.DiscoveryConfig
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -60,7 +60,7 @@ func (h *Handler) discoveryconfigCreate(w http.ResponseWriter, r *http.Request,
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -77,7 +77,7 @@ func (h *Handler) discoveryconfigCreate(w http.ResponseWriter, r *http.Request,
|
||||
}
|
||||
|
||||
// discoveryconfigUpdate updates the DiscoveryConfig based on its name
|
||||
func (h *Handler) discoveryconfigUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) discoveryconfigUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
dcName := p.ByName("name")
|
||||
if dcName == "" {
|
||||
return nil, trace.BadParameter("a discoveryconfig name is required")
|
||||
@@ -92,7 +92,7 @@ func (h *Handler) discoveryconfigUpdate(w http.ResponseWriter, r *http.Request,
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -118,13 +118,13 @@ func (h *Handler) discoveryconfigUpdate(w http.ResponseWriter, r *http.Request,
|
||||
}
|
||||
|
||||
// discoveryconfigDelete removes a DiscoveryConfig based on its name
|
||||
func (h *Handler) discoveryconfigDelete(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) discoveryconfigDelete(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
discoveryconfigName := p.ByName("name")
|
||||
if discoveryconfigName == "" {
|
||||
return nil, trace.BadParameter("a discoveryconfig name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -137,13 +137,13 @@ func (h *Handler) discoveryconfigDelete(w http.ResponseWriter, r *http.Request,
|
||||
}
|
||||
|
||||
// discoveryconfigGet returns a DiscoveryConfig based on its name
|
||||
func (h *Handler) discoveryconfigGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) discoveryconfigGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
discoveryconfigName := p.ByName("name")
|
||||
if discoveryconfigName == "" {
|
||||
return nil, trace.BadParameter("as discoveryconfig name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -157,8 +157,8 @@ func (h *Handler) discoveryconfigGet(w http.ResponseWriter, r *http.Request, p h
|
||||
}
|
||||
|
||||
// discoveryconfigList returns a page of DiscoveryConfigs
|
||||
func (h *Handler) discoveryconfigList(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) discoveryconfigList(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+4
-4
@@ -66,10 +66,10 @@ type fileTransferRequest struct {
|
||||
moderatedSessionID string
|
||||
}
|
||||
|
||||
func (h *Handler) transferFile(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) transferFile(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
query := r.URL.Query()
|
||||
req := fileTransferRequest{
|
||||
cluster: site.GetName(),
|
||||
cluster: cluster.GetName(),
|
||||
login: p.ByName("login"),
|
||||
serverID: p.ByName("server"),
|
||||
remoteLocation: query.Get("location"),
|
||||
@@ -98,7 +98,7 @@ func (h *Handler) transferFile(w http.ResponseWriter, r *http.Request, p httprou
|
||||
return nil, trace.BadParameter("fileTransferRequestId and moderatedSessionId must both be included in the same request.")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -160,7 +160,7 @@ func (h *Handler) transferFile(w http.ResponseWriter, r *http.Request, p httprou
|
||||
ctx = context.WithValue(ctx, sftp.ModeratedSessionID, req.moderatedSessionID)
|
||||
}
|
||||
|
||||
accessPoint, err := site.CachingAccessPoint()
|
||||
accessPoint, err := cluster.CachingAccessPoint()
|
||||
if err != nil {
|
||||
h.logger.DebugContext(r.Context(), "Unable to get auth access point", "error", err)
|
||||
return nil, trace.Wrap(err)
|
||||
|
||||
@@ -30,7 +30,7 @@ import (
|
||||
"github.com/gravitational/teleport/lib/web/ui"
|
||||
)
|
||||
|
||||
func (h *Handler) gitServerCreateOrUpsert(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) gitServerCreateOrUpsert(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var req *ui.CreateGitServerRequest
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -49,7 +49,7 @@ func (h *Handler) gitServerCreateOrUpsert(_ http.ResponseWriter, r *http.Request
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
userClient, err := sctx.GetUserClient(r.Context(), site)
|
||||
userClient, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -64,13 +64,13 @@ func (h *Handler) gitServerCreateOrUpsert(_ http.ResponseWriter, r *http.Request
|
||||
return created, trace.Wrap(err)
|
||||
}
|
||||
|
||||
func (h *Handler) gitServerGet(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) gitServerGet(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
name := p.ByName("name")
|
||||
if name == "" {
|
||||
return nil, trace.BadParameter("git server name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -79,16 +79,16 @@ func (h *Handler) gitServerGet(_ http.ResponseWriter, r *http.Request, p httprou
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
return ui.MakeGitServer(site.GetName(), gitServer, false), nil
|
||||
return ui.MakeGitServer(cluster.GetName(), gitServer, false), nil
|
||||
}
|
||||
|
||||
func (h *Handler) gitServerDelete(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) gitServerDelete(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
name := p.ByName("name")
|
||||
if name == "" {
|
||||
return nil, trace.BadParameter("git server name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+18
-18
@@ -46,7 +46,7 @@ import (
|
||||
)
|
||||
|
||||
// integrationsCreate creates an Integration
|
||||
func (h *Handler) integrationsCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) integrationsCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var req *ui.CreateIntegrationRequest
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -124,7 +124,7 @@ func (h *Handler) integrationsCreate(w http.ResponseWriter, r *http.Request, p h
|
||||
return nil, trace.BadParameter("subkind %q is not supported", req.SubKind)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -146,7 +146,7 @@ func (h *Handler) integrationsCreate(w http.ResponseWriter, r *http.Request, p h
|
||||
}
|
||||
|
||||
// integrationsUpdate updates the Integration based on its name
|
||||
func (h *Handler) integrationsUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) integrationsUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
integrationName := p.ByName("name")
|
||||
if integrationName == "" {
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
@@ -161,7 +161,7 @@ func (h *Handler) integrationsUpdate(w http.ResponseWriter, r *http.Request, p h
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -236,13 +236,13 @@ func (h *Handler) integrationsUpdate(w http.ResponseWriter, r *http.Request, p h
|
||||
}
|
||||
|
||||
// integrationsDelete removes an Integration based on its name
|
||||
func (h *Handler) integrationsDelete(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) integrationsDelete(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
integrationName := p.ByName("name_or_subkind")
|
||||
if integrationName == "" {
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -259,13 +259,13 @@ func (h *Handler) integrationsDelete(w http.ResponseWriter, r *http.Request, p h
|
||||
}
|
||||
|
||||
// integrationsGet returns an Integration based on its name
|
||||
func (h *Handler) integrationsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) integrationsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
integrationName := p.ByName("name")
|
||||
if integrationName == "" {
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -284,13 +284,13 @@ func (h *Handler) integrationsGet(w http.ResponseWriter, r *http.Request, p http
|
||||
}
|
||||
|
||||
// integrationStats returns the integration stats.
|
||||
func (h *Handler) integrationStats(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) integrationStats(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
integrationName := p.ByName("name")
|
||||
if integrationName == "" {
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -467,7 +467,7 @@ func rulesWithIntegration(dc *discoveryconfig.DiscoveryConfig, matcherType strin
|
||||
// startKey: indicator for pagination, should be the value of the last reponse's `nextItem`, or absent for a the starting page
|
||||
// resourceType: which resource type to return, one of ec2, eks, rds
|
||||
// regions: only rules for regions listed are returned (omit query to include all regions)
|
||||
func (h *Handler) integrationDiscoveryRules(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) integrationDiscoveryRules(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
integrationName := p.ByName("name")
|
||||
if integrationName == "" {
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
@@ -483,7 +483,7 @@ func (h *Handler) integrationDiscoveryRules(w http.ResponseWriter, r *http.Reque
|
||||
regionsFilter = nil
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -588,8 +588,8 @@ func collectAutoDiscoveryRulesFromDiscoveryConfig(dc *discoveryconfig.DiscoveryC
|
||||
}
|
||||
|
||||
// integrationsList returns a page of Integrations
|
||||
func (h *Handler) integrationsList(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) integrationsList(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -619,8 +619,8 @@ func (h *Handler) integrationsList(w http.ResponseWriter, r *http.Request, p htt
|
||||
}
|
||||
|
||||
// integrationsMsTeamsAppZipGet generates and returns the app.zip required for the MsTeams plugin with the given name.
|
||||
func (h *Handler) integrationsMsTeamsAppZipGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) integrationsMsTeamsAppZipGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -650,13 +650,13 @@ func (h *Handler) integrationsMsTeamsAppZipGet(w http.ResponseWriter, r *http.Re
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (h *Handler) integrationsExportCA(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) integrationsExportCA(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
integrationName := p.ByName("name")
|
||||
if integrationName == "" {
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -60,7 +60,7 @@ import (
|
||||
)
|
||||
|
||||
// awsOIDCListDatabases returns a list of databases using the ListDatabases action of the AWS OIDC Integration.
|
||||
func (h *Handler) awsOIDCListDatabases(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCListDatabases(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCListDatabasesRequest
|
||||
@@ -73,7 +73,7 @@ func (h *Handler) awsOIDCListDatabases(w http.ResponseWriter, r *http.Request, p
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -102,7 +102,7 @@ func (h *Handler) awsOIDCListDatabases(w http.ResponseWriter, r *http.Request, p
|
||||
}
|
||||
|
||||
// awsOIDCDeployService deploys a Discovery Service and a Database Service in Amazon ECS.
|
||||
func (h *Handler) awsOIDCDeployService(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCDeployService(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCDeployServiceRequest
|
||||
@@ -115,7 +115,7 @@ func (h *Handler) awsOIDCDeployService(w http.ResponseWriter, r *http.Request, p
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -195,7 +195,7 @@ func (h *Handler) awsOIDCDeployService(w http.ResponseWriter, r *http.Request, p
|
||||
}
|
||||
|
||||
// awsOIDCDeployDatabaseService deploys a Database Service in Amazon ECS.
|
||||
func (h *Handler) awsOIDCDeployDatabaseServices(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCDeployDatabaseServices(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCDeployDatabaseServiceRequest
|
||||
@@ -208,7 +208,7 @@ func (h *Handler) awsOIDCDeployDatabaseServices(w http.ResponseWriter, r *http.R
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -271,9 +271,9 @@ func (h *Handler) awsOIDCDeployDatabaseServices(w http.ResponseWriter, r *http.R
|
||||
}
|
||||
|
||||
// awsOIDCListDeployedDatabaseService lists the deployed Database Services in Amazon ECS.
|
||||
func (h *Handler) awsOIDCListDeployedDatabaseService(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCListDeployedDatabaseService(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -761,7 +761,7 @@ func (h *Handler) awsOIDCConfigureEKSIAM(w http.ResponseWriter, r *http.Request,
|
||||
|
||||
// awsOIDCEnrollEKSClusters enroll EKS clusters by installing teleport-kube-agent Helm chart on them.
|
||||
// v2 endpoint introduces "extraLabels" field.
|
||||
func (h *Handler) awsOIDCEnrollEKSClusters(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCEnrollEKSClusters(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCEnrollEKSClustersRequest
|
||||
@@ -769,7 +769,7 @@ func (h *Handler) awsOIDCEnrollEKSClusters(w http.ResponseWriter, r *http.Reques
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -818,7 +818,7 @@ func (h *Handler) awsOIDCEnrollEKSClusters(w http.ResponseWriter, r *http.Reques
|
||||
}
|
||||
|
||||
// awsOIDCListEKSClusters returns a list of EKS clusters using the ListEKSClusters action of the AWS OIDC integration.
|
||||
func (h *Handler) awsOIDCListEKSClusters(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCListEKSClusters(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCListEKSClustersRequest
|
||||
@@ -831,7 +831,7 @@ func (h *Handler) awsOIDCListEKSClusters(w http.ResponseWriter, r *http.Request,
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -852,7 +852,7 @@ func (h *Handler) awsOIDCListEKSClusters(w http.ResponseWriter, r *http.Request,
|
||||
}
|
||||
|
||||
// awsOIDCListSecurityGroups returns a list of VPC Security Groups using the ListSecurityGroups action of the AWS OIDC Integration.
|
||||
func (h *Handler) awsOIDCListSecurityGroups(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCListSecurityGroups(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCListSecurityGroupsRequest
|
||||
@@ -865,7 +865,7 @@ func (h *Handler) awsOIDCListSecurityGroups(w http.ResponseWriter, r *http.Reque
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -940,7 +940,7 @@ func awsOIDCSecurityGroupsRulesConverter(inRules []*integrationv1.SecurityGroupR
|
||||
// This api will return empty if we already have agents that can proxy the discovered databases.
|
||||
// Otherwise it will return with a map of VPC and its subnets where it's values are later used
|
||||
// to configure and deploy an agent (deploy an agent per unique VPC).
|
||||
func (h *Handler) awsOIDCRequiredDatabasesVPCS(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCRequiredDatabasesVPCS(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCRequiredVPCSRequest
|
||||
@@ -953,7 +953,7 @@ func (h *Handler) awsOIDCRequiredDatabasesVPCS(w http.ResponseWriter, r *http.Re
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -1058,7 +1058,7 @@ func awsOIDCRequiredVPCSHelper(ctx context.Context, clt client.GetResourcesClien
|
||||
|
||||
// awsOIDCCreateAWSAppAccess creates an AppServer that uses an AWS OIDC Integration for proxying access.
|
||||
// v2 endpoint introduces "labels" field
|
||||
func (h *Handler) awsOIDCCreateAWSAppAccess(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCCreateAWSAppAccess(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCCreateAWSAppAccessRequest
|
||||
@@ -1071,7 +1071,7 @@ func (h *Handler) awsOIDCCreateAWSAppAccess(w http.ResponseWriter, r *http.Reque
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -1137,7 +1137,7 @@ func (h *Handler) awsOIDCCreateAWSAppAccess(w http.ResponseWriter, r *http.Reque
|
||||
return ui.MakeApp(appServer.GetApp(), ui.MakeAppsConfig{
|
||||
LocalClusterName: h.auth.clusterName,
|
||||
LocalProxyDNSName: h.proxyDNSName(),
|
||||
AppClusterName: site.GetName(),
|
||||
AppClusterName: cluster.GetName(),
|
||||
AllowedAWSRolesLookup: allowedAWSRolesLookup,
|
||||
UserGroupLookup: getUserGroupLookup(),
|
||||
Logger: h.logger,
|
||||
@@ -1145,7 +1145,7 @@ func (h *Handler) awsOIDCCreateAWSAppAccess(w http.ResponseWriter, r *http.Reque
|
||||
}
|
||||
|
||||
// awsOIDCDeleteAWSAppAccess deletes the AWS AppServer created that uses the AWS OIDC Integration for proxying requests.
|
||||
func (h *Handler) awsOIDCDeleteAWSAppAccess(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCDeleteAWSAppAccess(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
subkind := p.ByName("name_or_subkind")
|
||||
@@ -1158,7 +1158,7 @@ func (h *Handler) awsOIDCDeleteAWSAppAccess(w http.ResponseWriter, r *http.Reque
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -1379,7 +1379,7 @@ func (h *Handler) awsAccessGraphOIDCSync(w http.ResponseWriter, r *http.Request,
|
||||
}
|
||||
|
||||
// awsOIDCListSubnets returns a list of VPC subnets using the ListSubnets action of the AWS OIDC Integration.
|
||||
func (h *Handler) awsOIDCListSubnets(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCListSubnets(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCListSubnetsRequest
|
||||
@@ -1392,7 +1392,7 @@ func (h *Handler) awsOIDCListSubnets(w http.ResponseWriter, r *http.Request, p h
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -1425,7 +1425,7 @@ func (h *Handler) awsOIDCListSubnets(w http.ResponseWriter, r *http.Request, p h
|
||||
// awsOIDCListDatabaseVPCs returns a list of VPCs using the ListVpcs action
|
||||
// of the AWS OIDC Integration, and includes a link to the ECS service if
|
||||
// a database service has been deployed for each VPC.
|
||||
func (h *Handler) awsOIDCListDatabaseVPCs(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCListDatabaseVPCs(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req ui.AWSOIDCListVPCsRequest
|
||||
@@ -1438,7 +1438,7 @@ func (h *Handler) awsOIDCListDatabaseVPCs(w http.ResponseWriter, r *http.Request
|
||||
return nil, trace.BadParameter("an integration name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -1552,7 +1552,7 @@ func getServiceURLs(dbServices []types.DatabaseService, accountID, region, telep
|
||||
// awsOIDCPing performs an health check for the integration.
|
||||
// If ARN is present in the request body, that's the ARN that will be used instead of using the one stored in the integration.
|
||||
// Returns meta information: account id and assumed the ARN for the IAM Role.
|
||||
func (h *Handler) awsOIDCPing(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsOIDCPing(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
integrationName := p.ByName("name")
|
||||
@@ -1565,7 +1565,7 @@ func (h *Handler) awsOIDCPing(w http.ResponseWriter, r *http.Request, p httprout
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -148,7 +148,7 @@ func (h *Handler) awsRolesAnywhereConfigureTrustAnchor(w http.ResponseWriter, r
|
||||
// It returns the caller identity and the number of AWS Roles Anywhere Profiles that are active.
|
||||
// If a trust anchor is provided in the body, it will be used to check the connection ignoring the integration.
|
||||
// Otherwise, the integration is used to check the connection.
|
||||
func (h *Handler) awsRolesAnywherePing(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsRolesAnywherePing(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
integrationName := p.ByName("name")
|
||||
@@ -161,7 +161,7 @@ func (h *Handler) awsRolesAnywherePing(w http.ResponseWriter, r *http.Request, p
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -204,7 +204,7 @@ func (h *Handler) awsRolesAnywherePing(w http.ResponseWriter, r *http.Request, p
|
||||
}
|
||||
|
||||
// awsRolesAnywhereListProfiles lists profiles Roles Anywhere Profiles accessible by the integration.
|
||||
func (h *Handler) awsRolesAnywhereListProfiles(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) awsRolesAnywhereListProfiles(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
integrationName := p.ByName("name")
|
||||
@@ -217,7 +217,7 @@ func (h *Handler) awsRolesAnywhereListProfiles(w http.ResponseWriter, r *http.Re
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+3
-3
@@ -437,7 +437,7 @@ func (h *Handler) joinKubernetesSession(
|
||||
sessionID string,
|
||||
mode types.SessionParticipantMode,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
ws *websocket.Conn,
|
||||
) error {
|
||||
h.logger.InfoContext(ctx, "Attempting to join kubernetes existing session",
|
||||
@@ -450,7 +450,7 @@ func (h *Handler) joinKubernetesSession(
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
@@ -499,7 +499,7 @@ func (h *Handler) joinKubernetesSession(
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
authAccessPoint, err := site.CachingAccessPoint()
|
||||
authAccessPoint, err := cluster.CachingAccessPoint()
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
|
||||
+18
-18
@@ -60,8 +60,8 @@ type CreateBotRequest struct {
|
||||
|
||||
// listBots returns a list of bots for a given cluster site. It does not leverage pagination from the UI. Due to the
|
||||
// nature of the bot:user relationship, pagination is not yet supported. This endpoint will return all bots.
|
||||
func (h *Handler) listBots(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) listBots(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -89,13 +89,13 @@ func (h *Handler) listBots(w http.ResponseWriter, r *http.Request, p httprouter.
|
||||
}
|
||||
|
||||
// createBot creates a bot
|
||||
func (h *Handler) createBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) createBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var req *CreateBotRequest
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -122,8 +122,8 @@ func (h *Handler) createBot(w http.ResponseWriter, r *http.Request, p httprouter
|
||||
return OK(), nil
|
||||
}
|
||||
|
||||
func (h *Handler) deleteBot(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) deleteBot(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -156,7 +156,7 @@ type CreateBotJoinTokenRequest struct {
|
||||
}
|
||||
|
||||
// createBotJoinToken creates a bot join token
|
||||
func (h *Handler) createBotJoinToken(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) createBotJoinToken(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var req *CreateBotJoinTokenRequest
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -166,7 +166,7 @@ func (h *Handler) createBotJoinToken(w http.ResponseWriter, r *http.Request, p h
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -198,13 +198,13 @@ func (h *Handler) createBotJoinToken(w http.ResponseWriter, r *http.Request, p h
|
||||
}
|
||||
|
||||
// getBot retrieves a bot by name
|
||||
func (h *Handler) getBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) getBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
botName := p.ByName("name")
|
||||
if botName == "" {
|
||||
return nil, trace.BadParameter("empty name")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -221,7 +221,7 @@ func (h *Handler) getBot(w http.ResponseWriter, r *http.Request, p httprouter.Pa
|
||||
// updateBot updates a bot with provided roles. The only supported change via this endpoint today is roles.
|
||||
// TODO(nicholasmarais1158) DELETE IN v20.0.0 - replaced by updateBotV2
|
||||
// MUST delete with related code found in `web/packages/teleport/src/services/bot/bot.ts`
|
||||
func (h *Handler) updateBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) updateBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var request updateBotRequest
|
||||
if err := httplib.ReadResourceJSON(r, &request); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -232,7 +232,7 @@ func (h *Handler) updateBot(w http.ResponseWriter, r *http.Request, p httprouter
|
||||
return nil, trace.BadParameter("empty name")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -267,7 +267,7 @@ type updateBotRequest struct {
|
||||
}
|
||||
|
||||
// updateBotV2 updates a bot with provided roles, traits and max_session_ttl.
|
||||
func (h *Handler) updateBotV2(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) updateBotV2(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var request updateBotRequestV2
|
||||
if err := httplib.ReadResourceJSON(r, &request); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -278,7 +278,7 @@ func (h *Handler) updateBotV2(w http.ResponseWriter, r *http.Request, p httprout
|
||||
return nil, trace.BadParameter("empty name")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -351,7 +351,7 @@ type updateBotRequestTrait struct {
|
||||
}
|
||||
|
||||
// getBotInstance retrieves a bot instance by id
|
||||
func (h *Handler) getBotInstance(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) getBotInstance(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
botName := p.ByName("name")
|
||||
instanceId := p.ByName("id")
|
||||
if botName == "" {
|
||||
@@ -361,7 +361,7 @@ func (h *Handler) getBotInstance(w http.ResponseWriter, r *http.Request, p httpr
|
||||
return nil, trace.BadParameter("empty id")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -390,8 +390,8 @@ type GetBotInstanceResponse struct {
|
||||
}
|
||||
|
||||
// listBotInstances returns a list of bot instances for a given cluster site.
|
||||
func (h *Handler) listBotInstances(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) listBotInstances(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+4
-4
@@ -601,13 +601,13 @@ type isMfaRequiredResponse struct {
|
||||
}
|
||||
|
||||
// isMFARequired is the [ClusterHandler] implementer for checking if MFA is required for a given target.
|
||||
func (h *Handler) isMFARequired(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) isMFARequired(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
var httpReq *IsMFARequiredRequest
|
||||
if err := httplib.ReadResourceJSON(r, &httpReq); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
required, err := h.checkMFARequired(r.Context(), httpReq, sctx, site)
|
||||
required, err := h.checkMFARequired(r.Context(), httpReq, sctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -616,13 +616,13 @@ func (h *Handler) isMFARequired(w http.ResponseWriter, r *http.Request, p httpro
|
||||
}
|
||||
|
||||
// checkMFARequired checks if MFA is required for the target specified in the [isMFARequiredRequest].
|
||||
func (h *Handler) checkMFARequired(ctx context.Context, req *IsMFARequiredRequest, sctx *SessionContext, site reversetunnelclient.RemoteSite) (bool, error) {
|
||||
func (h *Handler) checkMFARequired(ctx context.Context, req *IsMFARequiredRequest, sctx *SessionContext, cluster reversetunnelclient.Cluster) (bool, error) {
|
||||
protoReq, err := h.checkAndGetProtoRequest(ctx, sctx, req)
|
||||
if err != nil {
|
||||
return false, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return false, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+23
-23
@@ -35,8 +35,8 @@ import (
|
||||
)
|
||||
|
||||
// clusterKubesGet returns a list of kube clusters in a form the UI can present.
|
||||
func (h *Handler) clusterKubesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) clusterKubesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -64,7 +64,7 @@ func (h *Handler) clusterKubesGet(w http.ResponseWriter, r *http.Request, p http
|
||||
}
|
||||
|
||||
// clusterKubeResourcesGet returns supported requested kubernetes subresources eg: pods, namespaces, secrets etc.
|
||||
func (h *Handler) clusterKubeResourcesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) clusterKubeResourcesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
kind := r.URL.Query().Get("kind")
|
||||
kubeCluster := r.URL.Query().Get("kubeCluster")
|
||||
|
||||
@@ -85,7 +85,7 @@ func (h *Handler) clusterKubeResourcesGet(w http.ResponseWriter, r *http.Request
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
resp, err := listKubeResources(r.Context(), clt, r.URL.Query(), site.GetName(), kind)
|
||||
resp, err := listKubeResources(r.Context(), clt, r.URL.Query(), cluster.GetName(), kind)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -98,8 +98,8 @@ func (h *Handler) clusterKubeResourcesGet(w http.ResponseWriter, r *http.Request
|
||||
}
|
||||
|
||||
// clusterDatabasesGet returns a list of db servers in a form the UI can present.
|
||||
func (h *Handler) clusterDatabasesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) clusterDatabasesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -133,13 +133,13 @@ func (h *Handler) clusterDatabasesGet(w http.ResponseWriter, r *http.Request, p
|
||||
}
|
||||
|
||||
// clusterDatabaseGet returns a database in a form the UI can present.
|
||||
func (h *Handler) clusterDatabaseGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) clusterDatabaseGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
databaseName := p.ByName("database")
|
||||
if databaseName == "" {
|
||||
return nil, trace.BadParameter("database name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -171,8 +171,8 @@ func (h *Handler) clusterDatabaseGet(w http.ResponseWriter, r *http.Request, p h
|
||||
}
|
||||
|
||||
// clusterDatabaseServicesList returns a list of DatabaseServices (database agents) in a form the UI can present.
|
||||
func (h *Handler) clusterDatabaseServicesList(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := ctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) clusterDatabaseServicesList(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := ctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -195,8 +195,8 @@ func (h *Handler) clusterDatabaseServicesList(w http.ResponseWriter, r *http.Req
|
||||
}
|
||||
|
||||
// clusterDatabaseServersList returns a list of database servers in a form the UI can present.
|
||||
func (h *Handler) clusterDatabaseServersList(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := ctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) clusterDatabaseServersList(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := ctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -218,8 +218,8 @@ func (h *Handler) clusterDatabaseServersList(w http.ResponseWriter, r *http.Requ
|
||||
}
|
||||
|
||||
// clusterDesktopsGet returns a list of desktops in a form the UI can present.
|
||||
func (h *Handler) clusterDesktopsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) clusterDesktopsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -252,10 +252,10 @@ func (h *Handler) clusterDesktopsGet(w http.ResponseWriter, r *http.Request, p h
|
||||
}
|
||||
|
||||
// clusterDesktopServicesGet returns a list of desktop services in a form the UI can present.
|
||||
func (h *Handler) clusterDesktopServicesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) clusterDesktopServicesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
// Get a client to the Auth Server with the logged in user's identity. The
|
||||
// identity of the logged in user is used to fetch the list of desktop services.
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -278,8 +278,8 @@ func (h *Handler) clusterDesktopServicesGet(w http.ResponseWriter, r *http.Reque
|
||||
}
|
||||
|
||||
// getDesktopHandle returns a desktop.
|
||||
func (h *Handler) getDesktopHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) getDesktopHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -319,7 +319,7 @@ func (h *Handler) getDesktopHandle(w http.ResponseWriter, r *http.Request, p htt
|
||||
// Response body:
|
||||
//
|
||||
// {"active": bool}
|
||||
func (h *Handler) desktopIsActive(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) desktopIsActive(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
desktopName := p.ByName("desktopName")
|
||||
trackers, err := h.auth.proxyClient.GetActiveSessionTrackersWithFilter(r.Context(), &types.SessionTrackerFilter{
|
||||
Kind: string(types.WindowsDesktopSessionKind),
|
||||
@@ -332,7 +332,7 @@ func (h *Handler) desktopIsActive(w http.ResponseWriter, r *http.Request, p http
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -399,7 +399,7 @@ func (r *createNodeRequest) checkAndSetDefaults() error {
|
||||
}
|
||||
|
||||
// handleNodeCreate creates a Teleport Node.
|
||||
func (h *Handler) handleNodeCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) handleNodeCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
ctx := r.Context()
|
||||
|
||||
var req *createNodeRequest
|
||||
@@ -411,7 +411,7 @@ func (h *Handler) handleNodeCreate(w http.ResponseWriter, r *http.Request, p htt
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -458,5 +458,5 @@ func (h *Handler) handleNodeCreate(w http.ResponseWriter, r *http.Request, p htt
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
return webui.MakeServer(site.GetName(), server, logins, false /* requiresRequest */), nil
|
||||
return webui.MakeServer(cluster.GetName(), server, logins, false /* requiresRequest */), nil
|
||||
}
|
||||
|
||||
+37
-37
@@ -78,7 +78,7 @@ type SessionContext struct {
|
||||
// session.
|
||||
remoteClientCache
|
||||
// remoteClientGroup prevents duplicate requests to create remote clients
|
||||
// for a given site
|
||||
// for a given cluster
|
||||
remoteClientGroup singleflight.Group
|
||||
|
||||
// mu guards kubeGRPCServiceConn.
|
||||
@@ -116,8 +116,8 @@ type SessionContextConfig struct {
|
||||
// Session refers the web session created for the user.
|
||||
Session types.WebSession
|
||||
|
||||
// newRemoteClient is used by tests to override how remote clients are constructed to allow for fake sites
|
||||
newRemoteClient func(ctx context.Context, sessionContext *SessionContext, site reversetunnelclient.RemoteSite) (authclient.ClientI, error)
|
||||
// newRemoteClient is used by tests to override how remote clients are constructed to allow for fake clusters
|
||||
newRemoteClient func(ctx context.Context, sessionContext *SessionContext, cluster reversetunnelclient.Cluster) (authclient.ClientI, error)
|
||||
}
|
||||
|
||||
func (c *SessionContextConfig) CheckAndSetDefaults() error {
|
||||
@@ -228,36 +228,36 @@ func (c *SessionContext) GetClientConnection() *grpc.ClientConn {
|
||||
}
|
||||
|
||||
// GetUserClient will return an [authclient.ClientI] with the role of the user at
|
||||
// the requested site. If the site is local a client with the users local role
|
||||
// is returned. If the site is remote a client with the users remote role is
|
||||
// the requested cluster. If the cluster is local a client with the users local role
|
||||
// is returned. If the cluster is remote a client with the users remote role is
|
||||
// returned.
|
||||
func (c *SessionContext) GetUserClient(ctx context.Context, site reversetunnelclient.RemoteSite) (authclient.ClientI, error) {
|
||||
func (c *SessionContext) GetUserClient(ctx context.Context, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
|
||||
// if we're trying to access the local cluster, pass back the local client.
|
||||
if c.cfg.RootClusterName == site.GetName() {
|
||||
if c.cfg.RootClusterName == cluster.GetName() {
|
||||
return c.cfg.RootClient, nil
|
||||
}
|
||||
|
||||
// return the client for the requested remote site
|
||||
clt, err := c.remoteClient(ctx, site)
|
||||
// return the client for the requested remote cluster
|
||||
clt, err := c.remoteClient(ctx, cluster)
|
||||
return clt, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// remoteClient returns an [authclient.ClientI] with the role of the user at
|
||||
// the requested [site]. All remote clients are lazily created
|
||||
// the requested [cluster]. All remote clients are lazily created
|
||||
// when they are first requested and then cached. Subsequent requests
|
||||
// will return the previously created client to prevent having more than
|
||||
// a single [authclient.ClientI] per site for a user.
|
||||
// a single [authclient.ClientI] per cluster for a user.
|
||||
//
|
||||
// A [singleflight.Group] is leveraged to prevent duplicate requests for remote
|
||||
// clients at the same time to race.
|
||||
func (c *SessionContext) remoteClient(ctx context.Context, site reversetunnelclient.RemoteSite) (authclient.ClientI, error) {
|
||||
cltI, err, _ := c.remoteClientGroup.Do(site.GetName(), func() (any, error) {
|
||||
func (c *SessionContext) remoteClient(ctx context.Context, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
|
||||
cltI, err, _ := c.remoteClientGroup.Do(cluster.GetName(), func() (any, error) {
|
||||
// check if we already have a connection to this cluster
|
||||
if clt, ok := c.remoteClientCache.getRemoteClient(site); ok {
|
||||
if clt, ok := c.remoteClientCache.getRemoteClient(cluster); ok {
|
||||
return clt, nil
|
||||
}
|
||||
|
||||
rClt, err := c.cfg.newRemoteClient(ctx, c, site)
|
||||
rClt, err := c.cfg.newRemoteClient(ctx, c, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -265,10 +265,10 @@ func (c *SessionContext) remoteClient(ctx context.Context, site reversetunnelcli
|
||||
// we'll save the remote client in our session context so we don't have to
|
||||
// build a new connection next time. all remote clients will be closed when
|
||||
// the session context is closed.
|
||||
err = c.remoteClientCache.addRemoteClient(site, rClt)
|
||||
err = c.remoteClientCache.addRemoteClient(cluster, rClt)
|
||||
if err != nil {
|
||||
c.cfg.Log.InfoContext(ctx, "Failed closing stale remote client for site",
|
||||
"remote_site", site.GetName(),
|
||||
c.cfg.Log.InfoContext(ctx, "Failed closing stale remote client for cluster",
|
||||
"remote_cluster", cluster.GetName(),
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
@@ -289,8 +289,8 @@ func (c *SessionContext) remoteClient(ctx context.Context, site reversetunnelcli
|
||||
}
|
||||
|
||||
// newRemoteClient returns a client to a remote cluster with the role of current user.
|
||||
func newRemoteClient(ctx context.Context, sctx *SessionContext, site reversetunnelclient.RemoteSite) (authclient.ClientI, error) {
|
||||
clt, err := sctx.newRemoteTLSClient(ctx, site)
|
||||
func newRemoteClient(ctx context.Context, sctx *SessionContext, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
|
||||
clt, err := sctx.newRemoteTLSClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -306,7 +306,7 @@ func newRemoteClient(ctx context.Context, sctx *SessionContext, site reversetunn
|
||||
}
|
||||
|
||||
// clusterDialer returns DialContext function using cluster's dial function
|
||||
func clusterDialer(remoteCluster reversetunnelclient.RemoteSite, src, dst net.Addr) apiclient.ContextDialer {
|
||||
func clusterDialer(remoteCluster reversetunnelclient.Cluster, src, dst net.Addr) apiclient.ContextDialer {
|
||||
return apiclient.ContextDialerFunc(func(in context.Context, network, _ string) (net.Conn, error) {
|
||||
dialParams := reversetunnelclient.DialParams{
|
||||
From: src,
|
||||
@@ -400,7 +400,7 @@ func (c *SessionContext) ClientTLSConfig(ctx context.Context, clusterName ...str
|
||||
return tlsConfig, nil
|
||||
}
|
||||
|
||||
func (c *SessionContext) newRemoteTLSClient(ctx context.Context, cluster reversetunnelclient.RemoteSite) (authclient.ClientI, error) {
|
||||
func (c *SessionContext) newRemoteTLSClient(ctx context.Context, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
|
||||
tlsConfig, err := c.ClientTLSConfig(ctx, cluster.GetName())
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -615,7 +615,7 @@ func (c *SessionContext) expired(ctx context.Context) bool {
|
||||
// Give the session some time to linger so existing users of the context
|
||||
// have successfully disposed of them.
|
||||
// If we remove the session immediately, a stale copy might still use the
|
||||
// cached site clients.
|
||||
// cached cluster clients.
|
||||
// This is a cheaper way to avoid race without introducing object
|
||||
// reference counters.
|
||||
return c.cfg.Parent.clock.Since(expiry) > c.cfg.Parent.sessionLingeringThreshold
|
||||
@@ -1276,43 +1276,43 @@ func sessionKey(user, sessionID string) string {
|
||||
return user + sessionID
|
||||
}
|
||||
|
||||
// remoteClientCache stores remote clients keyed by site name while also keeping
|
||||
// track of the actual remote site associated with the client (in case the
|
||||
// remote site has changed). Safe for concurrent access. Closes all clients and
|
||||
// remoteClientCache stores remote clients keyed by cluster name while also keeping
|
||||
// track of the actual remote cluster associated with the client (in case the
|
||||
// remote cluster has changed). Safe for concurrent access. Closes all clients and
|
||||
// wipes the cache on Close.
|
||||
type remoteClientCache struct {
|
||||
sync.Mutex
|
||||
clients map[string]struct {
|
||||
authclient.ClientI
|
||||
reversetunnelclient.RemoteSite
|
||||
reversetunnelclient.Cluster
|
||||
}
|
||||
}
|
||||
|
||||
func (c *remoteClientCache) addRemoteClient(site reversetunnelclient.RemoteSite, remoteClient authclient.ClientI) error {
|
||||
func (c *remoteClientCache) addRemoteClient(cluster reversetunnelclient.Cluster, remoteClient authclient.ClientI) error {
|
||||
c.Lock()
|
||||
defer c.Unlock()
|
||||
if c.clients == nil {
|
||||
c.clients = make(map[string]struct {
|
||||
authclient.ClientI
|
||||
reversetunnelclient.RemoteSite
|
||||
reversetunnelclient.Cluster
|
||||
})
|
||||
}
|
||||
var err error
|
||||
if c.clients[site.GetName()].ClientI != nil {
|
||||
err = c.clients[site.GetName()].ClientI.Close()
|
||||
if c.clients[cluster.GetName()].ClientI != nil {
|
||||
err = c.clients[cluster.GetName()].ClientI.Close()
|
||||
}
|
||||
c.clients[site.GetName()] = struct {
|
||||
c.clients[cluster.GetName()] = struct {
|
||||
authclient.ClientI
|
||||
reversetunnelclient.RemoteSite
|
||||
}{remoteClient, site}
|
||||
reversetunnelclient.Cluster
|
||||
}{remoteClient, cluster}
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *remoteClientCache) getRemoteClient(site reversetunnelclient.RemoteSite) (authclient.ClientI, bool) {
|
||||
func (c *remoteClientCache) getRemoteClient(cluster reversetunnelclient.Cluster) (authclient.ClientI, bool) {
|
||||
c.Lock()
|
||||
defer c.Unlock()
|
||||
remoteClt, ok := c.clients[site.GetName()]
|
||||
return remoteClt.ClientI, ok && remoteClt.RemoteSite == site
|
||||
remoteClt, ok := c.clients[cluster.GetName()]
|
||||
return remoteClt.ClientI, ok && remoteClt.Cluster == cluster
|
||||
}
|
||||
|
||||
func (c *remoteClientCache) Close() error {
|
||||
|
||||
+11
-11
@@ -45,9 +45,9 @@ func TestRemoteClientCache(t *testing.T) {
|
||||
var openCount atomic.Int32
|
||||
cache := remoteClientCache{}
|
||||
|
||||
sa1 := newMockRemoteSite("a")
|
||||
sa2 := newMockRemoteSite("a")
|
||||
sb := newMockRemoteSite("b")
|
||||
sa1 := newMockCluster("a")
|
||||
sa2 := newMockCluster("a")
|
||||
sb := newMockCluster("b")
|
||||
|
||||
err1 := errors.New("c1")
|
||||
err2 := errors.New("c2")
|
||||
@@ -68,16 +68,16 @@ func TestRemoteClientCache(t *testing.T) {
|
||||
require.Zero(t, openCount.Load())
|
||||
}
|
||||
|
||||
func newMockRemoteSite(name string) reversetunnelclient.RemoteSite {
|
||||
return &mockRemoteSite{name: name}
|
||||
func newMockCluster(name string) reversetunnelclient.Cluster {
|
||||
return &mockCluster{name: name}
|
||||
}
|
||||
|
||||
type mockRemoteSite struct {
|
||||
reversetunnelclient.RemoteSite
|
||||
type mockCluster struct {
|
||||
reversetunnelclient.Cluster
|
||||
name string
|
||||
}
|
||||
|
||||
func (m *mockRemoteSite) GetName() string {
|
||||
func (m *mockCluster) GetName() string {
|
||||
return m.name
|
||||
}
|
||||
|
||||
@@ -109,14 +109,14 @@ func TestGetUserClient(t *testing.T) {
|
||||
sctx := SessionContext{
|
||||
cfg: SessionContextConfig{
|
||||
RootClusterName: "local",
|
||||
newRemoteClient: func(ctx context.Context, sessionContext *SessionContext, site reversetunnelclient.RemoteSite) (authclient.ClientI, error) {
|
||||
newRemoteClient: func(ctx context.Context, sessionContext *SessionContext, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
|
||||
return newMockClientI(&openCount, nil), nil
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
localSite := &mockRemoteSite{name: "local"}
|
||||
remoteSite := &mockRemoteSite{name: "remote"}
|
||||
localSite := &mockCluster{name: "local"}
|
||||
remoteSite := &mockCluster{name: "remote"}
|
||||
|
||||
// getting a client for the local site should return
|
||||
// the RootClient from SessionContextConfig
|
||||
|
||||
+1
-1
@@ -53,7 +53,7 @@ Should be equivalent to running:
|
||||
|
||||
This endpoint returns a tar.gz compressed archive containing the required files to setup mTLS for the database.
|
||||
*/
|
||||
func (h *Handler) signDatabaseCertificate(w http.ResponseWriter, r *http.Request, p httprouter.Params, site reversetunnelclient.RemoteSite, token types.ProvisionToken) (any, error) {
|
||||
func (h *Handler) signDatabaseCertificate(w http.ResponseWriter, r *http.Request, p httprouter.Params, cluster reversetunnelclient.Cluster, token types.ProvisionToken) (any, error) {
|
||||
if !token.GetRoles().Include(types.RoleDatabase) {
|
||||
return nil, trace.AccessDenied("required '%s' role was not provided by the token", types.RoleDatabase)
|
||||
}
|
||||
|
||||
@@ -59,7 +59,7 @@ func (h *Handler) sessionLengthHandle(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
) (any, error) {
|
||||
sID := p.ByName("sid")
|
||||
if sID == "" {
|
||||
@@ -69,7 +69,7 @@ func (h *Handler) sessionLengthHandle(
|
||||
ctx, cancel := context.WithCancel(r.Context())
|
||||
defer cancel()
|
||||
|
||||
clt, err := sctx.GetUserClient(ctx, site)
|
||||
clt, err := sctx.GetUserClient(ctx, cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -106,14 +106,14 @@ func (h *Handler) ttyPlaybackHandle(
|
||||
r *http.Request,
|
||||
p httprouter.Params,
|
||||
sctx *SessionContext,
|
||||
site reversetunnelclient.RemoteSite,
|
||||
cluster reversetunnelclient.Cluster,
|
||||
) (any, error) {
|
||||
sID := p.ByName("sid")
|
||||
if sID == "" {
|
||||
return nil, trace.BadParameter("missing session ID in request URL")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
+10
-10
@@ -52,16 +52,16 @@ type Cluster struct {
|
||||
}
|
||||
|
||||
// NewClusters creates a slice of Cluster's, containing data about each cluster.
|
||||
func NewClusters(remoteClusters []reversetunnelclient.RemoteSite) ([]Cluster, error) {
|
||||
func NewClusters(remoteClusters []reversetunnelclient.Cluster) ([]Cluster, error) {
|
||||
clusters := make([]Cluster, 0, len(remoteClusters))
|
||||
for _, site := range remoteClusters {
|
||||
for _, cluster := range remoteClusters {
|
||||
// Other fields such as node count, url, and proxy/auth versions are not set
|
||||
// because each cluster will need to make network calls to retrieve information
|
||||
// which does not scale well (ie: 1k clusters, each request will take seconds).
|
||||
cluster := &Cluster{
|
||||
Name: site.GetName(),
|
||||
LastConnected: site.GetLastConnected(),
|
||||
Status: site.GetStatus(),
|
||||
Name: cluster.GetName(),
|
||||
LastConnected: cluster.GetLastConnected(),
|
||||
Status: cluster.GetStatus(),
|
||||
}
|
||||
|
||||
clusters = append(clusters, *cluster)
|
||||
@@ -89,8 +89,8 @@ func NewClustersFromRemote(remoteClusters []types.RemoteCluster) ([]Cluster, err
|
||||
}
|
||||
|
||||
// GetClusterDetails retrieves and sets details about a cluster
|
||||
func GetClusterDetails(ctx context.Context, site reversetunnelclient.RemoteSite, opts ...services.MarshalOption) (*Cluster, error) {
|
||||
clt, err := site.CachingAccessPoint()
|
||||
func GetClusterDetails(ctx context.Context, cluster reversetunnelclient.Cluster, opts ...services.MarshalOption) (*Cluster, error) {
|
||||
clt, err := cluster.CachingAccessPoint()
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -119,9 +119,9 @@ func GetClusterDetails(ctx context.Context, site reversetunnelclient.RemoteSite,
|
||||
}
|
||||
|
||||
return &Cluster{
|
||||
Name: site.GetName(),
|
||||
LastConnected: site.GetLastConnected(),
|
||||
Status: site.GetStatus(),
|
||||
Name: cluster.GetName(),
|
||||
LastConnected: cluster.GetLastConnected(),
|
||||
Status: cluster.GetStatus(),
|
||||
PublicURL: proxyHost,
|
||||
AuthVersion: authVersion,
|
||||
|
||||
|
||||
+11
-11
@@ -89,7 +89,7 @@ func BenchmarkGetClusterDetails(b *testing.B) {
|
||||
insertServers(ctx, b, svc, types.KindProxy, proxyCount)
|
||||
insertServers(ctx, b, svc, types.KindAuthServer, authCount)
|
||||
|
||||
site := &mockRemoteSite{
|
||||
site := &mockCluster{
|
||||
accessPoint: &mockAccessPoint{
|
||||
presence: svc,
|
||||
},
|
||||
@@ -137,34 +137,34 @@ func insertServers(ctx context.Context, b *testing.B, svc services.Presence, kin
|
||||
}
|
||||
}
|
||||
|
||||
func benchmarkGetClusterDetails(ctx context.Context, b *testing.B, site reversetunnelclient.RemoteSite, nodes int, opts ...services.MarshalOption) {
|
||||
var cluster *Cluster
|
||||
func benchmarkGetClusterDetails(ctx context.Context, b *testing.B, cluster reversetunnelclient.Cluster, nodes int, opts ...services.MarshalOption) {
|
||||
var got *Cluster
|
||||
var err error
|
||||
for b.Loop() {
|
||||
cluster, err = GetClusterDetails(ctx, site, opts...)
|
||||
got, err = GetClusterDetails(ctx, cluster, opts...)
|
||||
require.NoError(b, err)
|
||||
}
|
||||
require.NotNil(b, cluster)
|
||||
require.NotNil(b, got)
|
||||
}
|
||||
|
||||
type mockRemoteSite struct {
|
||||
reversetunnelclient.RemoteSite
|
||||
type mockCluster struct {
|
||||
reversetunnelclient.Cluster
|
||||
accessPoint authclient.ProxyAccessPoint
|
||||
}
|
||||
|
||||
func (m *mockRemoteSite) CachingAccessPoint() (authclient.RemoteProxyAccessPoint, error) {
|
||||
func (m *mockCluster) CachingAccessPoint() (authclient.RemoteProxyAccessPoint, error) {
|
||||
return m.accessPoint, nil
|
||||
}
|
||||
|
||||
func (m *mockRemoteSite) GetName() string {
|
||||
func (m *mockCluster) GetName() string {
|
||||
return clusterName
|
||||
}
|
||||
|
||||
func (m *mockRemoteSite) GetLastConnected() time.Time {
|
||||
func (m *mockCluster) GetLastConnected() time.Time {
|
||||
return time.Now()
|
||||
}
|
||||
|
||||
func (m *mockRemoteSite) GetStatus() string {
|
||||
func (m *mockCluster) GetStatus() string {
|
||||
return teleport.RemoteClusterStatusOnline
|
||||
}
|
||||
|
||||
|
||||
@@ -33,10 +33,10 @@ import (
|
||||
"github.com/gravitational/teleport/lib/web/ui"
|
||||
)
|
||||
|
||||
func (h *Handler) getUserGroups(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) getUserGroups(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
// Get a client to the Auth Server with the logged in user's identity. The
|
||||
// identity of the logged in user is used to fetch the list of nodes.
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -89,8 +89,8 @@ type UserPreferencesResponse struct {
|
||||
KeyboardLayout uint32 `json:"keyboardLayout"`
|
||||
}
|
||||
|
||||
func (h *Handler) getUserClusterPreferences(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
authClient, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) getUserClusterPreferences(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
authClient, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -104,14 +104,14 @@ func (h *Handler) getUserClusterPreferences(_ http.ResponseWriter, r *http.Reque
|
||||
}
|
||||
|
||||
// updateUserClusterPreferences is a handler for PUT /webapi/user/preferences.
|
||||
func (h *Handler) updateUserClusterPreferences(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) updateUserClusterPreferences(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
req := UserPreferencesResponse{}
|
||||
|
||||
if err := httplib.ReadResourceJSON(r, &req); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
authClient, err := sctx.GetUserClient(r.Context(), site)
|
||||
authClient, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
@@ -32,7 +32,7 @@ import (
|
||||
)
|
||||
|
||||
// userTaskStateUpdate updates the state of a User Task.
|
||||
func (h *Handler) userTaskStateUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) userTaskStateUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
userTaskName := p.ByName("name")
|
||||
if userTaskName == "" {
|
||||
return nil, trace.BadParameter("a user task name is required")
|
||||
@@ -47,7 +47,7 @@ func (h *Handler) userTaskStateUpdate(w http.ResponseWriter, r *http.Request, p
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -67,13 +67,13 @@ func (h *Handler) userTaskStateUpdate(w http.ResponseWriter, r *http.Request, p
|
||||
}
|
||||
|
||||
// userTaskGet returns a User Task based on its name
|
||||
func (h *Handler) userTaskGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
func (h *Handler) userTaskGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
userTaskName := p.ByName("name")
|
||||
if userTaskName == "" {
|
||||
return nil, trace.BadParameter("a user task name is required")
|
||||
}
|
||||
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
@@ -95,8 +95,8 @@ func (h *Handler) userTaskGet(w http.ResponseWriter, r *http.Request, p httprout
|
||||
//
|
||||
// It returns a list of user tasks with the base attributes (common among all user tasks).
|
||||
// To get a detailed UserTask use the single resource endpoint, ie, usertask/<resource's name>.
|
||||
func (h *Handler) userTaskListByIntegration(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, site reversetunnelclient.RemoteSite) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), site)
|
||||
func (h *Handler) userTaskListByIntegration(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
|
||||
clt, err := sctx.GetUserClient(r.Context(), cluster)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user