Paginated rpcs - Replace GetNodes with ListNodes (#7415)

This commit is contained in:
Brian Joerger
2021-07-01 14:03:23 -07:00
committed by GitHub
parent 91f5c88554
commit 20a734ec5d
18 changed files with 1522 additions and 373 deletions
+52 -7
View File
@@ -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
+8 -8
View File
@@ -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()) })
}
File diff suppressed because it is too large Load Diff
+21
View File
@@ -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.
+5
View File
@@ -68,3 +68,8 @@ func EnhancedEvents() []string {
constants.EnhancedRecordingNetwork,
}
}
const (
// DefaultChunkSize is the default chunk size for paginated endpoints.
DefaultChunkSize = 1000
)
+1
View File
@@ -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"`
+1
View File
@@ -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;
+3
View File
@@ -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)
+31
View File
@@ -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...)
+44
View File
@@ -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)
+57
View File
@@ -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))
}
+24
View File
@@ -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)
+172
View File
@@ -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
View File
@@ -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
+10
View File
@@ -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()
+42
View File
@@ -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) {
+116
View File
@@ -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))
})
}
+3
View File
@@ -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