mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Paginated rpcs - Replace GetNodes with ListNodes (#7415)
This commit is contained in:
+52
-7
@@ -1261,21 +1261,66 @@ func (c *Client) GetNode(ctx context.Context, namespace, name string) (types.Ser
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// GetNodes returns a list of nodes by namespace.
|
||||
// Nodes that the user doesn't have access to are filtered out.
|
||||
// GetNodes returns a complete list of nodes that the user has access to in the given namespace.
|
||||
func (c *Client) GetNodes(ctx context.Context, namespace string) ([]types.Server, error) {
|
||||
if namespace == "" {
|
||||
return nil, trace.BadParameter("missing parameter namespace")
|
||||
}
|
||||
resp, err := c.grpc.GetNodes(ctx, &types.ResourcesInNamespaceRequest{Namespace: namespace}, c.callOpts...)
|
||||
if err != nil {
|
||||
return nil, trail.FromGRPC(err)
|
||||
|
||||
// Retrieve the complete list of nodes in chunks.
|
||||
var (
|
||||
nodes []types.Server
|
||||
startKey string
|
||||
chunkSize = defaults.DefaultChunkSize
|
||||
)
|
||||
for {
|
||||
resp, nextKey, err := c.ListNodes(ctx, namespace, chunkSize, startKey)
|
||||
if trace.IsLimitExceeded(err) {
|
||||
// Cut chunkSize in half if gRPC max message size is exceeded.
|
||||
chunkSize = chunkSize / 2
|
||||
// This is an extremely unlikely scenario, but better to cover it anyways.
|
||||
if chunkSize == 0 {
|
||||
return nil, trace.Wrap(trail.FromGRPC(err), "Node is too large to retrieve over gRPC (over 4MiB).")
|
||||
}
|
||||
continue
|
||||
} else if err != nil {
|
||||
return nil, trail.FromGRPC(err)
|
||||
}
|
||||
|
||||
nodes = append(nodes, resp...)
|
||||
startKey = nextKey
|
||||
if startKey == "" {
|
||||
return nodes, nil
|
||||
}
|
||||
}
|
||||
nodes := make([]types.Server, len(resp.Servers))
|
||||
}
|
||||
|
||||
// ListNodes returns a paginated list of nodes that the user has access to in the given namespace.
|
||||
// nextKey can be used as startKey in another call to ListNodes to retrieve the next page of nodes.
|
||||
// ListNodes will return a trace.LimitExceeded error if the page of nodes retrieved exceeds 4MiB.
|
||||
func (c *Client) ListNodes(ctx context.Context, namespace string, limit int, startKey string) (nodes []types.Server, nextKey string, err error) {
|
||||
if namespace == "" {
|
||||
return nil, "", trace.BadParameter("missing parameter namespace")
|
||||
}
|
||||
if limit <= 0 {
|
||||
return nil, "", trace.BadParameter("nonpositive parameter limit")
|
||||
}
|
||||
|
||||
resp, err := c.grpc.ListNodes(ctx, &proto.ListNodesRequest{
|
||||
Namespace: namespace,
|
||||
Limit: int32(limit),
|
||||
StartKey: startKey,
|
||||
}, c.callOpts...)
|
||||
if err != nil {
|
||||
return nil, "", trail.FromGRPC(err)
|
||||
}
|
||||
|
||||
nodes = make([]types.Server, len(resp.Servers))
|
||||
for i, node := range resp.Servers {
|
||||
nodes[i] = node
|
||||
}
|
||||
return nodes, nil
|
||||
|
||||
return nodes, resp.NextKey, nil
|
||||
}
|
||||
|
||||
// UpsertNode is used by SSH servers to report their presence
|
||||
|
||||
@@ -46,6 +46,14 @@ func newMockServer() *mockServer {
|
||||
return m
|
||||
}
|
||||
|
||||
// startMockServer starts a new mock server. Parallel tests cannot use the same addr.
|
||||
func startMockServer(t *testing.T, addr string) {
|
||||
l, err := net.Listen("tcp", addr)
|
||||
require.NoError(t, err)
|
||||
go newMockServer().grpc.Serve(l)
|
||||
t.Cleanup(func() { require.NoError(t, l.Close()) })
|
||||
}
|
||||
|
||||
func (m *mockServer) Ping(ctx context.Context, req *proto.PingRequest) (*proto.PingResponse, error) {
|
||||
return &proto.PingResponse{}, nil
|
||||
}
|
||||
@@ -202,11 +210,3 @@ func TestWaitForConnectionReady(t *testing.T) {
|
||||
require.NoError(t, clt.GetConnection().Close())
|
||||
require.Error(t, clt.waitForConnectionReady(ctx))
|
||||
}
|
||||
|
||||
// startMockServer starts a new mock server. Parallel tests cannot use the same addr.
|
||||
func startMockServer(t *testing.T, addr string) {
|
||||
l, err := net.Listen("tcp", addr)
|
||||
require.NoError(t, err)
|
||||
go newMockServer().grpc.Serve(l)
|
||||
t.Cleanup(func() { require.NoError(t, l.Close()) })
|
||||
}
|
||||
|
||||
+917
-356
File diff suppressed because it is too large
Load Diff
@@ -867,6 +867,24 @@ message Events {
|
||||
string LastKey = 2;
|
||||
}
|
||||
|
||||
message ListNodesRequest {
|
||||
// Namespace is the namespace of resources.
|
||||
string Namespace = 1;
|
||||
// Limit is the maximum amount of nodes to retrieve.
|
||||
int32 Limit = 2;
|
||||
// StartKey is used to start listing nodes from a specific spot. This should
|
||||
// be set to the previous NextKey value if using pagination, or left empty.
|
||||
string StartKey = 3;
|
||||
}
|
||||
|
||||
message ListNodesResponse {
|
||||
// Servers is a list of servers.
|
||||
repeated types.ServerV2 Servers = 1;
|
||||
// NextKey is the next Key to use as StartKey in a ListNodesRequest to continue
|
||||
// retrieving pages of nodes. If NextKey is empty, there are no more pages.
|
||||
string NextKey = 2;
|
||||
}
|
||||
|
||||
// AuthService is authentication/authorization service implementation
|
||||
service AuthService {
|
||||
// SendKeepAlives allows node to send a stream of keep alive requests
|
||||
@@ -877,7 +895,10 @@ service AuthService {
|
||||
// GetNode retrieves a node described by the given request.
|
||||
rpc GetNode(types.ResourceInNamespaceRequest) returns (types.ServerV2);
|
||||
// GetNodes retrieves all nodes.
|
||||
// DELETE IN 8.0.0 in favor of ListNodes
|
||||
rpc GetNodes(types.ResourcesInNamespaceRequest) returns (types.ServerV2List);
|
||||
// ListNodes retrieves a paginated list of nodes.
|
||||
rpc ListNodes(ListNodesRequest) returns (ListNodesResponse);
|
||||
// UpsertNode upserts a node in a backend.
|
||||
rpc UpsertNode(types.ServerV2) returns (types.KeepAlive);
|
||||
// DeleteNode deletes an existing node in a backend described by the given request.
|
||||
|
||||
@@ -68,3 +68,8 @@ func EnhancedEvents() []string {
|
||||
constants.EnhancedRecordingNetwork,
|
||||
}
|
||||
}
|
||||
|
||||
const (
|
||||
// DefaultChunkSize is the default chunk size for paginated endpoints.
|
||||
DefaultChunkSize = 1000
|
||||
)
|
||||
|
||||
@@ -714,6 +714,7 @@ func (m *ServerV2) XXX_DiscardUnknown() {
|
||||
var xxx_messageInfo_ServerV2 proto.InternalMessageInfo
|
||||
|
||||
// ServerV2List is a list of servers.
|
||||
// DELETE IN 8.0.0 only used in deprecated GetNodes rpc
|
||||
type ServerV2List struct {
|
||||
// Servers is a list of servers.
|
||||
Servers []*ServerV2 `protobuf:"bytes,1,rep,name=Servers,proto3" json:"Servers,omitempty"`
|
||||
|
||||
@@ -217,6 +217,7 @@ message ServerV2 {
|
||||
}
|
||||
|
||||
// ServerV2List is a list of servers.
|
||||
// DELETE IN 8.0.0 only used in deprecated GetNodes rpc
|
||||
message ServerV2List {
|
||||
// Servers is a list of servers.
|
||||
repeated ServerV2 Servers = 1;
|
||||
|
||||
@@ -98,6 +98,9 @@ type ReadAccessPoint interface {
|
||||
// GetNodes returns a list of registered servers for this cluster.
|
||||
GetNodes(ctx context.Context, namespace string, opts ...services.MarshalOption) ([]types.Server, error)
|
||||
|
||||
// ListNodes returns a paginated list of registered servers for this cluster.
|
||||
ListNodes(ctx context.Context, namespace string, limit int, startKey string) (nodes []types.Server, nextKey string, err error)
|
||||
|
||||
// GetProxies returns a list of proxy servers registered in the cluster
|
||||
GetProxies() ([]types.Server, error)
|
||||
|
||||
|
||||
@@ -2122,6 +2122,37 @@ func (a *Server) GetNodes(ctx context.Context, namespace string, opts ...service
|
||||
return a.GetCache().GetNodes(ctx, namespace, opts...)
|
||||
}
|
||||
|
||||
// ListNodes is a part of auth.AccessPoint implementation
|
||||
func (a *Server) ListNodes(ctx context.Context, namespace string, limit int, startKey string) ([]types.Server, string, error) {
|
||||
return a.GetCache().ListNodes(ctx, namespace, limit, startKey)
|
||||
}
|
||||
|
||||
// NodePageFunc is a function to run on each page iterated over.
|
||||
type NodePageFunc func(next []types.Server) (stop bool, err error)
|
||||
|
||||
// IterateNodePages can be used to iterate over pages of nodes.
|
||||
func (a *Server) IterateNodePages(ctx context.Context, namespace string, limit int, startKey string, f NodePageFunc) (string, error) {
|
||||
for {
|
||||
nextPage, nextKey, err := a.ListNodes(ctx, namespace, limit, startKey)
|
||||
if err != nil {
|
||||
return "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
stop, err := f(nextPage)
|
||||
if err != nil {
|
||||
return "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
// Iterator stopped before end of pages or
|
||||
// there are no more pages, return nextKey
|
||||
if stop || nextKey == "" {
|
||||
return nextKey, nil
|
||||
}
|
||||
|
||||
startKey = nextKey
|
||||
}
|
||||
}
|
||||
|
||||
// GetReverseTunnels is a part of auth.AccessPoint implementation
|
||||
func (a *Server) GetReverseTunnels(opts ...services.MarshalOption) ([]types.ReverseTunnel, error) {
|
||||
return a.GetCache().GetReverseTunnels(opts...)
|
||||
|
||||
@@ -30,6 +30,7 @@ import (
|
||||
"github.com/gravitational/teleport/api/types/wrappers"
|
||||
apiutils "github.com/gravitational/teleport/api/utils"
|
||||
"github.com/gravitational/teleport/lib/auth/u2f"
|
||||
"github.com/gravitational/teleport/lib/backend"
|
||||
"github.com/gravitational/teleport/lib/defaults"
|
||||
"github.com/gravitational/teleport/lib/events"
|
||||
"github.com/gravitational/teleport/lib/modules"
|
||||
@@ -661,6 +662,49 @@ func (a *ServerWithRoles) GetNodes(ctx context.Context, namespace string, opts .
|
||||
return filteredNodes, nil
|
||||
}
|
||||
|
||||
// ListNodes returns a paginated list of nodes filtered by user access.
|
||||
func (a *ServerWithRoles) ListNodes(ctx context.Context, namespace string, limit int, startKey string) (page []types.Server, nextKey string, err error) {
|
||||
if err := a.action(namespace, types.KindNode, types.VerbList); err != nil {
|
||||
return nil, "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
return a.filterAndListNodes(ctx, namespace, limit, startKey)
|
||||
}
|
||||
|
||||
func (a *ServerWithRoles) filterAndListNodes(ctx context.Context, namespace string, limit int, startKey string) (page []types.Server, nextKey string, err error) {
|
||||
if limit <= 0 {
|
||||
return nil, "", trace.BadParameter("nonpositive parameter limit")
|
||||
}
|
||||
|
||||
page = make([]types.Server, 0, limit)
|
||||
nextKey, err = a.authServer.IterateNodePages(ctx, namespace, limit, startKey, func(nextPage []types.Server) (bool, error) {
|
||||
// Retrieve and filter pages of nodes until we can fill a page or run out of nodes.
|
||||
filteredPage, err := a.filterNodes(nextPage)
|
||||
if err != nil {
|
||||
return false, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// We have more than enough nodes to fill the page, cut it to size.
|
||||
if len(filteredPage) > limit-len(page) {
|
||||
filteredPage = filteredPage[:limit-len(page)]
|
||||
}
|
||||
|
||||
// Add filteredPage and break out of iterator if the page is now full.
|
||||
page = append(page, filteredPage...)
|
||||
return len(page) == limit, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
// Filled a page, reset nextKey in case the last node was cut out.
|
||||
if len(page) == limit {
|
||||
nextKey = backend.NextPaginationKey(page[len(page)-1])
|
||||
}
|
||||
|
||||
return page, nextKey, nil
|
||||
}
|
||||
|
||||
func (a *ServerWithRoles) UpsertAuthServer(s types.Server) error {
|
||||
if err := a.action(apidefaults.Namespace, types.KindAuthServer, types.VerbCreate); err != nil {
|
||||
return trace.Wrap(err)
|
||||
|
||||
@@ -24,11 +24,13 @@ import (
|
||||
|
||||
"github.com/gravitational/teleport/api/client/proto"
|
||||
"github.com/gravitational/teleport/api/constants"
|
||||
"github.com/gravitational/teleport/api/defaults"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
"github.com/google/go-cmp/cmp/cmpopts"
|
||||
"github.com/pborman/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
@@ -239,3 +241,58 @@ func resourceDiff(res1, res2 types.Resource) string {
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID", "Namespace"),
|
||||
cmpopts.EquateEmpty())
|
||||
}
|
||||
|
||||
// TestListNodes users can retrieve nodes with the appropriate permissions.
|
||||
func TestListNodes(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
srv := newTestTLSServer(t)
|
||||
|
||||
// Create test nodes.
|
||||
for i := 0; i < 10; i++ {
|
||||
name := uuid.New()
|
||||
node, err := types.NewServerWithLabels(
|
||||
name,
|
||||
types.KindNode,
|
||||
types.ServerSpecV2{},
|
||||
map[string]string{"name": name},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = srv.Auth().UpsertNode(ctx, node)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
testNodes, err := srv.Auth().GetNodes(ctx, defaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
|
||||
// create user, role, and client
|
||||
username := "user"
|
||||
user, role, err := CreateUserAndRole(srv.Auth(), username, nil)
|
||||
require.NoError(t, err)
|
||||
identity := TestUser(user.GetName())
|
||||
clt, err := srv.NewClient(identity)
|
||||
require.NoError(t, err)
|
||||
|
||||
// permit user to list all nodes
|
||||
role.SetNodeLabels(types.Allow, types.Labels{types.Wildcard: {types.Wildcard}})
|
||||
require.NoError(t, srv.Auth().UpsertRole(ctx, role))
|
||||
|
||||
// listing nodes 0-4 should list first 5 nodes
|
||||
nodes, _, err := clt.ListNodes(ctx, defaults.Namespace, 5, "")
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 5, len(nodes))
|
||||
expectedNodes := testNodes[:5]
|
||||
require.Empty(t, cmp.Diff(expectedNodes, nodes))
|
||||
|
||||
// remove permission for third node
|
||||
role.SetNodeLabels(types.Deny, types.Labels{"name": {testNodes[3].GetName()}})
|
||||
require.NoError(t, srv.Auth().UpsertRole(ctx, role))
|
||||
|
||||
// listing nodes 0-4 should skip the third node and add the fifth to the end.
|
||||
nodes, _, err = clt.ListNodes(ctx, defaults.Namespace, 5, "")
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 5, len(nodes))
|
||||
expectedNodes = append(testNodes[:3], testNodes[4:6]...)
|
||||
require.Empty(t, cmp.Diff(expectedNodes, nodes))
|
||||
}
|
||||
|
||||
@@ -2424,6 +2424,7 @@ func (g *GRPCServer) GetNode(ctx context.Context, req *types.ResourceInNamespace
|
||||
}
|
||||
|
||||
// GetNodes retrieves all nodes in the given namespace.
|
||||
// DELETE IN 8.0.0 in favor of ListNodes
|
||||
func (g *GRPCServer) GetNodes(ctx context.Context, req *types.ResourcesInNamespaceRequest) (*types.ServerV2List, error) {
|
||||
auth, err := g.authenticate(ctx)
|
||||
if err != nil {
|
||||
@@ -2445,6 +2446,29 @@ func (g *GRPCServer) GetNodes(ctx context.Context, req *types.ResourcesInNamespa
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ListNodes retrieves a paginated list of nodes in the given namespace.
|
||||
func (g *GRPCServer) ListNodes(ctx context.Context, req *proto.ListNodesRequest) (*proto.ListNodesResponse, error) {
|
||||
auth, err := g.authenticate(ctx)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
ns, nextKey, err := auth.ServerWithRoles.ListNodes(ctx, req.Namespace, int(req.Limit), req.StartKey)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
serversV2 := make([]*types.ServerV2, len(ns))
|
||||
for i, t := range ns {
|
||||
var ok bool
|
||||
if serversV2[i], ok = t.(*types.ServerV2); !ok {
|
||||
return nil, trace.Errorf("encountered unexpected node type: %T", t)
|
||||
}
|
||||
}
|
||||
return &proto.ListNodesResponse{
|
||||
Servers: serversV2,
|
||||
NextKey: nextKey,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// UpsertNode upserts a node.
|
||||
func (g *GRPCServer) UpsertNode(ctx context.Context, node *types.ServerV2) (*types.KeepAlive, error) {
|
||||
auth, err := g.authenticate(ctx)
|
||||
|
||||
@@ -36,13 +36,17 @@ import (
|
||||
"github.com/gravitational/teleport"
|
||||
"github.com/gravitational/teleport/api/client/proto"
|
||||
"github.com/gravitational/teleport/api/constants"
|
||||
"github.com/gravitational/teleport/api/defaults"
|
||||
apidefaults "github.com/gravitational/teleport/api/defaults"
|
||||
"github.com/gravitational/teleport/api/metadata"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/api/utils/sshutils"
|
||||
"github.com/gravitational/teleport/lib/auth/mocku2f"
|
||||
"github.com/gravitational/teleport/lib/auth/u2f"
|
||||
"github.com/gravitational/teleport/lib/backend"
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
"github.com/gravitational/teleport/lib/tlsca"
|
||||
"github.com/gravitational/trace"
|
||||
)
|
||||
|
||||
func TestMFADeviceManagement(t *testing.T) {
|
||||
@@ -1206,3 +1210,171 @@ func TestSessionRecordingConfigOriginDynamic(t *testing.T) {
|
||||
|
||||
testOriginDynamicStored(t, setWithOrigin, getStored)
|
||||
}
|
||||
|
||||
func TestNodesCRUD(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
srv := newTestTLSServer(t)
|
||||
|
||||
clt, err := srv.NewClient(TestAdmin())
|
||||
require.NoError(t, err)
|
||||
|
||||
// node1 and node2 will be added to default namespace
|
||||
node1, err := types.NewServerWithLabels("node1", types.KindNode, types.ServerSpecV2{}, map[string]string{
|
||||
// Artificially make a node ~ 3MB to force ListNodes
|
||||
// to fail when retrieving more than one at a time.
|
||||
"label": string(make([]byte, 3000000)),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
node2, err := types.NewServerWithLabels("node2", types.KindNode, types.ServerSpecV2{}, map[string]string{
|
||||
// Artificially make a node ~ 3MB to force ListNodes
|
||||
// to fail when retrieving more than one at a time.
|
||||
"label": string(make([]byte, 3000000)),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// add largeNode to special namespace. largeNode is too big to send over gRPC.
|
||||
largeNode, err := types.NewServerWithLabels(
|
||||
"the_big_one",
|
||||
types.KindNode,
|
||||
types.ServerSpecV2{},
|
||||
map[string]string{
|
||||
// Artificially make a node ~ 5MB to force
|
||||
// ListNodes to fail regardless of chunk size.
|
||||
"label": string(make([]byte, 5000000)),
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
largeNodeNamespace := "the_big_one"
|
||||
largeNode.SetNamespace(largeNodeNamespace)
|
||||
_, err = srv.Auth().UpsertNode(ctx, largeNode)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("CreateNode", func(t *testing.T) {
|
||||
// Initially expect no nodes to be returned.
|
||||
nodes, err := clt.GetNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 0, len(nodes))
|
||||
|
||||
// Create nodes
|
||||
_, err = clt.UpsertNode(ctx, node1)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = clt.UpsertNode(ctx, node2)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Fail to create largeNode
|
||||
_, err = clt.UpsertNode(ctx, largeNode)
|
||||
require.IsType(t, &trace.LimitExceededError{}, err.(*trace.TraceErr).OrigError())
|
||||
})
|
||||
|
||||
// Run NodeGetters in nested subtests to allow parallelization.
|
||||
t.Run("NodeGetters", func(t *testing.T) {
|
||||
t.Run("List Nodes", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
// list nodes one at a time, last page should be empty
|
||||
nodes, nextKey, err := clt.ListNodes(ctx, apidefaults.Namespace, 1, "")
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, len(nodes))
|
||||
require.Empty(t, cmp.Diff([]types.Server{node1}, nodes,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
require.EqualValues(t, backend.NextPaginationKey(node1), nextKey)
|
||||
|
||||
nodes, nextKey, err = clt.ListNodes(ctx, apidefaults.Namespace, 1, nextKey)
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, len(nodes))
|
||||
require.Empty(t, cmp.Diff([]types.Server{node2}, nodes,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
require.EqualValues(t, backend.NextPaginationKey(node2), nextKey)
|
||||
|
||||
nodes, nextKey, err = clt.ListNodes(ctx, apidefaults.Namespace, 1, nextKey)
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 0, len(nodes))
|
||||
require.EqualValues(t, "", nextKey)
|
||||
|
||||
// ListNodes should fail if namespace isn't provided
|
||||
_, _, err = clt.ListNodes(ctx, "", 1, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
// ListNodes should fail if limit is nonpositive
|
||||
_, _, err = clt.ListNodes(ctx, apidefaults.Namespace, 0, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
_, _, err = clt.ListNodes(ctx, apidefaults.Namespace, -1, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
// ListNodes should return a limit exceeded error when exceeding gRPC message size limit.
|
||||
_, _, err = clt.ListNodes(ctx, defaults.Namespace, 2, "")
|
||||
require.IsType(t, &trace.LimitExceededError{}, err.(*trace.TraceErr).OrigError())
|
||||
})
|
||||
t.Run("GetNodes", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Get all nodes, transparently handle limit exceeded errors
|
||||
nodes, err := clt.GetNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, len(nodes), 2)
|
||||
require.Empty(t, cmp.Diff([]types.Server{node1, node2}, nodes,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
|
||||
// GetNodes should fail if namespace isn't provided
|
||||
_, err = clt.GetNodes(ctx, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
// GetNodes should return a limit exceeded error when a single
|
||||
// node is larger than the gRPC message size limit.
|
||||
_, err = clt.GetNodes(ctx, largeNodeNamespace)
|
||||
require.IsType(t, &trace.LimitExceededError{}, err.(*trace.TraceErr).OrigError())
|
||||
})
|
||||
t.Run("GetNode", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Get Node
|
||||
node, err := clt.GetNode(ctx, apidefaults.Namespace, "node1")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cmp.Diff(node1, node,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
|
||||
// GetNode should fail if node name isn't provided
|
||||
_, err = clt.GetNode(ctx, apidefaults.Namespace, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
// GetNode should fail if namespace isn't provided
|
||||
_, err = clt.GetNode(ctx, "", "node1")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("DeleteNode", func(t *testing.T) {
|
||||
// Make sure can't delete with empty namespace or name.
|
||||
err = clt.DeleteNode(ctx, apidefaults.Namespace, "")
|
||||
require.Error(t, err)
|
||||
require.IsType(t, trace.BadParameter(""), err)
|
||||
|
||||
err = clt.DeleteNode(ctx, "", node1.GetName())
|
||||
require.Error(t, err)
|
||||
require.IsType(t, trace.BadParameter(""), err)
|
||||
|
||||
// Delete node.
|
||||
err = clt.DeleteNode(ctx, apidefaults.Namespace, node1.GetName())
|
||||
require.NoError(t, err)
|
||||
|
||||
// Expect node not found
|
||||
_, err := clt.GetNode(ctx, apidefaults.Namespace, "node1")
|
||||
require.IsType(t, trace.NotFound(""), err)
|
||||
})
|
||||
|
||||
t.Run("DeleteAllNodes", func(t *testing.T) {
|
||||
// Make sure can't delete with empty namespace.
|
||||
err = clt.DeleteAllNodes(ctx, "")
|
||||
require.Error(t, err)
|
||||
require.IsType(t, trace.BadParameter(""), err)
|
||||
|
||||
// Delete nodes
|
||||
err = clt.DeleteAllNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Now expect no nodes to be returned.
|
||||
nodes, err := clt.GetNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 0, len(nodes))
|
||||
})
|
||||
}
|
||||
|
||||
+15
-2
@@ -26,6 +26,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
|
||||
"github.com/jonboulle/clockwork"
|
||||
)
|
||||
|
||||
@@ -212,8 +213,10 @@ func (p Params) GetString(key string) string {
|
||||
// NoLimit specifies no limits
|
||||
const NoLimit = 0
|
||||
|
||||
// RangeEnd returns end of the range for given key
|
||||
func RangeEnd(key []byte) []byte {
|
||||
// nextKey returns the next possible key.
|
||||
// If used with a key prefix, this will return
|
||||
// the end of the range for that key prefix.
|
||||
func nextKey(key []byte) []byte {
|
||||
end := make([]byte, len(key))
|
||||
copy(end, key)
|
||||
for i := len(end) - 1; i >= 0; i-- {
|
||||
@@ -231,6 +234,16 @@ var (
|
||||
noEnd = []byte{0}
|
||||
)
|
||||
|
||||
// RangeEnd returns end of the range for given key.
|
||||
func RangeEnd(key []byte) []byte {
|
||||
return nextKey(key)
|
||||
}
|
||||
|
||||
// NextPaginationKey returns the next pagination key.
|
||||
func NextPaginationKey(r types.Resource) string {
|
||||
return string(nextKey([]byte(r.GetName())))
|
||||
}
|
||||
|
||||
// Items is a sortable list of backend items
|
||||
type Items []Item
|
||||
|
||||
|
||||
Vendored
+10
@@ -1148,6 +1148,16 @@ func (c *Cache) GetNodes(ctx context.Context, namespace string, opts ...services
|
||||
return rg.presence.GetNodes(ctx, namespace, opts...)
|
||||
}
|
||||
|
||||
// ListNodes is a part of auth.AccessPoint implementation
|
||||
func (c *Cache) ListNodes(ctx context.Context, namespace string, limit int, startKey string) ([]types.Server, string, error) {
|
||||
rg, err := c.read()
|
||||
if err != nil {
|
||||
return nil, "", trace.Wrap(err)
|
||||
}
|
||||
defer rg.Release()
|
||||
return rg.presence.ListNodes(ctx, namespace, limit, startKey)
|
||||
}
|
||||
|
||||
// GetAuthServers returns a list of registered servers
|
||||
func (c *Cache) GetAuthServers() ([]types.Server, error) {
|
||||
rg, err := c.read()
|
||||
|
||||
@@ -251,6 +251,48 @@ func (s *PresenceService) GetNodes(ctx context.Context, namespace string, opts .
|
||||
return servers, nil
|
||||
}
|
||||
|
||||
// ListNodes returns a paginated list of registered servers.
|
||||
// StartKey is a resource name, which is the suffix of its key.
|
||||
func (s *PresenceService) ListNodes(ctx context.Context, namespace string, limit int, startKey string) (page []types.Server, nextKey string, err error) {
|
||||
if namespace == "" {
|
||||
return nil, "", trace.BadParameter("missing namespace value")
|
||||
}
|
||||
if limit <= 0 {
|
||||
return nil, "", trace.BadParameter("nonpositive limit value")
|
||||
}
|
||||
|
||||
// Get all items in the bucket within the given range.
|
||||
rangeStart := backend.Key(nodesPrefix, namespace, startKey)
|
||||
keyPrefix := backend.Key(nodesPrefix, namespace)
|
||||
rangeEnd := backend.RangeEnd(keyPrefix)
|
||||
result, err := s.GetRange(ctx, rangeStart, rangeEnd, limit)
|
||||
if err != nil {
|
||||
return nil, "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
// Marshal values into a []services.Server slice.
|
||||
servers := make([]types.Server, len(result.Items))
|
||||
for i, item := range result.Items {
|
||||
server, err := services.UnmarshalServer(
|
||||
item.Value,
|
||||
types.KindNode,
|
||||
services.WithResourceID(item.ID),
|
||||
services.WithExpires(item.Expires),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, "", trace.Wrap(err)
|
||||
}
|
||||
servers[i] = server
|
||||
}
|
||||
|
||||
// If a full page was filled, set nextKey using the last node.
|
||||
if len(result.Items) == limit {
|
||||
nextKey = backend.NextPaginationKey(servers[len(servers)-1])
|
||||
}
|
||||
|
||||
return servers, nextKey, nil
|
||||
}
|
||||
|
||||
// UpsertNode registers node presence, permanently if TTL is 0 or for the
|
||||
// specified duration with second resolution if it's >= 1 second.
|
||||
func (s *PresenceService) UpsertNode(ctx context.Context, server types.Server) (*types.KeepAlive, error) {
|
||||
|
||||
@@ -21,6 +21,8 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
"github.com/google/go-cmp/cmp/cmpopts"
|
||||
"github.com/jonboulle/clockwork"
|
||||
"github.com/pborman/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -191,3 +193,117 @@ func TestDatabaseServersCRUD(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 0, len(out))
|
||||
}
|
||||
|
||||
func TestNodeCRUD(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
lite, err := lite.NewWithConfig(ctx, lite.Config{Path: t.TempDir()})
|
||||
require.NoError(t, err)
|
||||
|
||||
presence := NewPresenceService(lite)
|
||||
|
||||
node1, err := types.NewServerWithLabels("node1", types.KindNode, types.ServerSpecV2{}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
node2, err := types.NewServerWithLabels("node2", types.KindNode, types.ServerSpecV2{}, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("CreateNode", func(t *testing.T) {
|
||||
// Initially expect no nodes to be returned.
|
||||
nodes, err := presence.GetNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 0, len(nodes))
|
||||
|
||||
// Create nodes
|
||||
_, err = presence.UpsertNode(ctx, node1)
|
||||
require.NoError(t, err)
|
||||
_, err = presence.UpsertNode(ctx, node2)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
// Run NodeGetters in nested subtests to allow parallelization.
|
||||
t.Run("NodeGetters", func(t *testing.T) {
|
||||
t.Run("List Nodes", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
// list nodes one at a time, last page should be empty
|
||||
nodes, nextKey, err := presence.ListNodes(ctx, apidefaults.Namespace, 1, "")
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, len(nodes))
|
||||
require.Empty(t, cmp.Diff([]types.Server{node1}, nodes,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
require.EqualValues(t, backend.NextPaginationKey(node1), nextKey)
|
||||
|
||||
nodes, nextKey, err = presence.ListNodes(ctx, apidefaults.Namespace, 1, nextKey)
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 1, len(nodes))
|
||||
require.Empty(t, cmp.Diff([]types.Server{node2}, nodes,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
require.EqualValues(t, backend.NextPaginationKey(node2), nextKey)
|
||||
|
||||
nodes, nextKey, err = presence.ListNodes(ctx, apidefaults.Namespace, 1, nextKey)
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 0, len(nodes))
|
||||
require.EqualValues(t, "", nextKey)
|
||||
|
||||
// ListNodes should fail if namespace isn't provided
|
||||
_, _, err = presence.ListNodes(ctx, "", 1, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
// ListNodes should fail if limit is nonpositive
|
||||
_, _, err = presence.ListNodes(ctx, apidefaults.Namespace, 0, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
_, _, err = presence.ListNodes(ctx, apidefaults.Namespace, -1, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
})
|
||||
t.Run("GetNodes", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Get all nodes, transparently handle limit exceeded errors
|
||||
nodes, err := presence.GetNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, len(nodes), 2)
|
||||
require.Empty(t, cmp.Diff([]types.Server{node1, node2}, nodes,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
|
||||
// GetNodes should fail if namespace isn't provided
|
||||
_, err = presence.GetNodes(ctx, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
})
|
||||
t.Run("GetNode", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Get Node
|
||||
node, err := presence.GetNode(ctx, apidefaults.Namespace, "node1")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cmp.Diff(node1, node,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "ID")))
|
||||
|
||||
// GetNode should fail if node name isn't provided
|
||||
_, err = presence.GetNode(ctx, apidefaults.Namespace, "")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
|
||||
// GetNode should fail if namespace isn't provided
|
||||
_, err = presence.GetNode(ctx, "", "node1")
|
||||
require.IsType(t, &trace.BadParameterError{}, err.(*trace.TraceErr).OrigError())
|
||||
})
|
||||
})
|
||||
|
||||
t.Run("DeleteNode", func(t *testing.T) {
|
||||
// Delete node.
|
||||
err = presence.DeleteNode(ctx, apidefaults.Namespace, node1.GetName())
|
||||
require.NoError(t, err)
|
||||
|
||||
// Expect node not found
|
||||
_, err := presence.GetNode(ctx, apidefaults.Namespace, "node1")
|
||||
require.IsType(t, trace.NotFound(""), err)
|
||||
})
|
||||
|
||||
t.Run("DeleteAllNodes", func(t *testing.T) {
|
||||
// Delete nodes
|
||||
err = presence.DeleteAllNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Now expect no nodes to be returned.
|
||||
nodes, err := presence.GetNodes(ctx, apidefaults.Namespace)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 0, len(nodes))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -40,6 +40,9 @@ type Presence interface {
|
||||
// GetNodes returns a list of registered servers.
|
||||
GetNodes(ctx context.Context, namespace string, opts ...MarshalOption) ([]types.Server, error)
|
||||
|
||||
// ListNodes returns a paginated list of registered servers.
|
||||
ListNodes(ctx context.Context, namespace string, limit int, startKey string) (nodes []types.Server, nextKey string, err error)
|
||||
|
||||
// DeleteAllNodes deletes all nodes in a namespace.
|
||||
DeleteAllNodes(ctx context.Context, namespace string) error
|
||||
|
||||
|
||||
Reference in New Issue
Block a user