mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-19 11:00:37 +08:00
[api] Add ListDatabases (#57805)
This commits deperecates GetDatabases, which should be replaced by a paginated counterpart: ListDatabases The commit also adds TakeWhile for conditional iter stream reads. Towards: gravitational/teleport.e/issues/6759 Changelog: Add paginated API ListDatabases, deprecate GetDatabases
This commit is contained in:
@@ -3536,7 +3536,9 @@ func (c *Client) GetDatabase(ctx context.Context, name string) (types.Database,
|
||||
//
|
||||
// For a full list of registered databases that are served by a database
|
||||
// service, use GetDatabaseServers instead.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
func (c *Client) GetDatabases(ctx context.Context) ([]types.Database, error) {
|
||||
//nolint:staticcheck // TODO(okraport): deprecated, to be removed in v21
|
||||
items, err := c.grpc.GetDatabases(ctx, &emptypb.Empty{})
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -3548,6 +3550,58 @@ func (c *Client) GetDatabases(ctx context.Context) ([]types.Database, error) {
|
||||
return databases, nil
|
||||
}
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
//
|
||||
// Note that database resources here refers to "dynamically-added" databases
|
||||
// such as databases created by `tctl create`, the discovery service, or the
|
||||
// CreateDatabase API. Databases discovered by the database agent (legacy
|
||||
// discovery flow using `database_service.aws/database_service.azure`) and
|
||||
// static databases defined in the `database_service.databases` section of the
|
||||
// service YAML configuration are not collected in this API.
|
||||
func (c *Client) ListDatabases(ctx context.Context, limit int, start string) ([]types.Database, string, error) {
|
||||
resp, err := c.grpc.ListDatabases(ctx, &proto.ListDatabasesRequest{
|
||||
PageSize: int32(limit),
|
||||
PageToken: start,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, "", trace.Wrap(err)
|
||||
}
|
||||
databases := make([]types.Database, len(resp.Databases))
|
||||
for i := range resp.Databases {
|
||||
databases[i] = resp.Databases[i]
|
||||
}
|
||||
return databases, resp.NextPageToken, nil
|
||||
}
|
||||
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
func (c *Client) RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error] {
|
||||
return func(yield func(types.Database, error) bool) {
|
||||
for {
|
||||
databases, next, err := c.ListDatabases(ctx, 0, start)
|
||||
if err != nil {
|
||||
yield(nil, err)
|
||||
return
|
||||
}
|
||||
|
||||
for _, db := range databases {
|
||||
if end != "" && db.GetName() >= end {
|
||||
return
|
||||
}
|
||||
|
||||
if !yield(db, nil) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if next == "" {
|
||||
return
|
||||
}
|
||||
|
||||
start = next
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteDatabase deletes specified database resource.
|
||||
func (c *Client) DeleteDatabase(ctx context.Context, name string) error {
|
||||
_, err := c.grpc.DeleteDatabase(ctx, &types.ResourceRequest{Name: name})
|
||||
|
||||
+1512
-1044
File diff suppressed because it is too large
Load Diff
@@ -227,6 +227,7 @@ const (
|
||||
AuthService_DeleteApp_FullMethodName = "/proto.AuthService/DeleteApp"
|
||||
AuthService_DeleteAllApps_FullMethodName = "/proto.AuthService/DeleteAllApps"
|
||||
AuthService_GetDatabases_FullMethodName = "/proto.AuthService/GetDatabases"
|
||||
AuthService_ListDatabases_FullMethodName = "/proto.AuthService/ListDatabases"
|
||||
AuthService_GetDatabase_FullMethodName = "/proto.AuthService/GetDatabase"
|
||||
AuthService_CreateDatabase_FullMethodName = "/proto.AuthService/CreateDatabase"
|
||||
AuthService_UpdateDatabase_FullMethodName = "/proto.AuthService/UpdateDatabase"
|
||||
@@ -806,8 +807,11 @@ type AuthServiceClient interface {
|
||||
DeleteApp(ctx context.Context, in *types.ResourceRequest, opts ...grpc.CallOption) (*emptypb.Empty, error)
|
||||
// DeleteAllApps removes all application resources.
|
||||
DeleteAllApps(ctx context.Context, in *emptypb.Empty, opts ...grpc.CallOption) (*emptypb.Empty, error)
|
||||
// Deprecated: Do not use.
|
||||
// GetDatabases returns all registered databases.
|
||||
GetDatabases(ctx context.Context, in *emptypb.Empty, opts ...grpc.CallOption) (*types.DatabaseV3List, error)
|
||||
// ListDatabases returns a page of registered databases.
|
||||
ListDatabases(ctx context.Context, in *ListDatabasesRequest, opts ...grpc.CallOption) (*ListDatabasesResponse, error)
|
||||
// GetDatabase returns a database by name.
|
||||
GetDatabase(ctx context.Context, in *types.ResourceRequest, opts ...grpc.CallOption) (*types.DatabaseV3, error)
|
||||
// CreateDatabase creates a new database resource.
|
||||
@@ -3066,6 +3070,7 @@ func (c *authServiceClient) DeleteAllApps(ctx context.Context, in *emptypb.Empty
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Deprecated: Do not use.
|
||||
func (c *authServiceClient) GetDatabases(ctx context.Context, in *emptypb.Empty, opts ...grpc.CallOption) (*types.DatabaseV3List, error) {
|
||||
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
|
||||
out := new(types.DatabaseV3List)
|
||||
@@ -3076,6 +3081,16 @@ func (c *authServiceClient) GetDatabases(ctx context.Context, in *emptypb.Empty,
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *authServiceClient) ListDatabases(ctx context.Context, in *ListDatabasesRequest, opts ...grpc.CallOption) (*ListDatabasesResponse, error) {
|
||||
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
|
||||
out := new(ListDatabasesResponse)
|
||||
err := c.cc.Invoke(ctx, AuthService_ListDatabases_FullMethodName, in, out, cOpts...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *authServiceClient) GetDatabase(ctx context.Context, in *types.ResourceRequest, opts ...grpc.CallOption) (*types.DatabaseV3, error) {
|
||||
cOpts := append([]grpc.CallOption{grpc.StaticMethod()}, opts...)
|
||||
out := new(types.DatabaseV3)
|
||||
@@ -4310,8 +4325,11 @@ type AuthServiceServer interface {
|
||||
DeleteApp(context.Context, *types.ResourceRequest) (*emptypb.Empty, error)
|
||||
// DeleteAllApps removes all application resources.
|
||||
DeleteAllApps(context.Context, *emptypb.Empty) (*emptypb.Empty, error)
|
||||
// Deprecated: Do not use.
|
||||
// GetDatabases returns all registered databases.
|
||||
GetDatabases(context.Context, *emptypb.Empty) (*types.DatabaseV3List, error)
|
||||
// ListDatabases returns a page of registered databases.
|
||||
ListDatabases(context.Context, *ListDatabasesRequest) (*ListDatabasesResponse, error)
|
||||
// GetDatabase returns a database by name.
|
||||
GetDatabase(context.Context, *types.ResourceRequest) (*types.DatabaseV3, error)
|
||||
// CreateDatabase creates a new database resource.
|
||||
@@ -5100,6 +5118,9 @@ func (UnimplementedAuthServiceServer) DeleteAllApps(context.Context, *emptypb.Em
|
||||
func (UnimplementedAuthServiceServer) GetDatabases(context.Context, *emptypb.Empty) (*types.DatabaseV3List, error) {
|
||||
return nil, status.Errorf(codes.Unimplemented, "method GetDatabases not implemented")
|
||||
}
|
||||
func (UnimplementedAuthServiceServer) ListDatabases(context.Context, *ListDatabasesRequest) (*ListDatabasesResponse, error) {
|
||||
return nil, status.Errorf(codes.Unimplemented, "method ListDatabases not implemented")
|
||||
}
|
||||
func (UnimplementedAuthServiceServer) GetDatabase(context.Context, *types.ResourceRequest) (*types.DatabaseV3, error) {
|
||||
return nil, status.Errorf(codes.Unimplemented, "method GetDatabase not implemented")
|
||||
}
|
||||
@@ -8627,6 +8648,24 @@ func _AuthService_GetDatabases_Handler(srv interface{}, ctx context.Context, dec
|
||||
return interceptor(ctx, in, info, handler)
|
||||
}
|
||||
|
||||
func _AuthService_ListDatabases_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) {
|
||||
in := new(ListDatabasesRequest)
|
||||
if err := dec(in); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if interceptor == nil {
|
||||
return srv.(AuthServiceServer).ListDatabases(ctx, in)
|
||||
}
|
||||
info := &grpc.UnaryServerInfo{
|
||||
Server: srv,
|
||||
FullMethod: AuthService_ListDatabases_FullMethodName,
|
||||
}
|
||||
handler := func(ctx context.Context, req interface{}) (interface{}, error) {
|
||||
return srv.(AuthServiceServer).ListDatabases(ctx, req.(*ListDatabasesRequest))
|
||||
}
|
||||
return interceptor(ctx, in, info, handler)
|
||||
}
|
||||
|
||||
func _AuthService_GetDatabase_Handler(srv interface{}, ctx context.Context, dec func(interface{}) error, interceptor grpc.UnaryServerInterceptor) (interface{}, error) {
|
||||
in := new(types.ResourceRequest)
|
||||
if err := dec(in); err != nil {
|
||||
@@ -10619,6 +10658,10 @@ var AuthService_ServiceDesc = grpc.ServiceDesc{
|
||||
MethodName: "GetDatabases",
|
||||
Handler: _AuthService_GetDatabases_Handler,
|
||||
},
|
||||
{
|
||||
MethodName: "ListDatabases",
|
||||
Handler: _AuthService_ListDatabases_Handler,
|
||||
},
|
||||
{
|
||||
MethodName: "GetDatabase",
|
||||
Handler: _AuthService_GetDatabase_Handler,
|
||||
|
||||
@@ -1633,6 +1633,22 @@ message ListAppsResponse {
|
||||
string next_key = 2;
|
||||
}
|
||||
|
||||
message ListDatabasesRequest {
|
||||
// The maximum number of items to return.
|
||||
// The server may impose a different page size at its discretion.
|
||||
int32 page_size = 1;
|
||||
// The next_page_token value returned from a previous List request, if any.
|
||||
string page_token = 2;
|
||||
}
|
||||
|
||||
message ListDatabasesResponse {
|
||||
// a list of databases.
|
||||
repeated types.DatabaseV3 databases = 1;
|
||||
// Token to retrieve the next page of results, or empty if there are no
|
||||
// more results in the list.
|
||||
string next_page_token = 2;
|
||||
}
|
||||
|
||||
// GetWindowsDesktopServicesResponse contains all registered Windows desktop services.
|
||||
message GetWindowsDesktopServicesResponse {
|
||||
// Services is a list of Windows desktop services.
|
||||
@@ -3338,7 +3354,11 @@ service AuthService {
|
||||
rpc DeleteAllApps(google.protobuf.Empty) returns (google.protobuf.Empty);
|
||||
|
||||
// GetDatabases returns all registered databases.
|
||||
rpc GetDatabases(google.protobuf.Empty) returns (types.DatabaseV3List);
|
||||
rpc GetDatabases(google.protobuf.Empty) returns (types.DatabaseV3List) {
|
||||
option deprecated = true;
|
||||
}
|
||||
// ListDatabases returns a page of registered databases.
|
||||
rpc ListDatabases(ListDatabasesRequest) returns (ListDatabasesResponse);
|
||||
// GetDatabase returns a database by name.
|
||||
rpc GetDatabase(types.ResourceRequest) returns (types.DatabaseV3);
|
||||
// CreateDatabase creates a new database resource.
|
||||
|
||||
@@ -130,7 +130,6 @@ func newDefaultConfig() *Config {
|
||||
"proto.AuthService.GetAlertAcks": {},
|
||||
"proto.AuthService.GetApps": {},
|
||||
"proto.AuthService.GetClusterAlerts": {},
|
||||
"proto.AuthService.GetDatabases": {},
|
||||
"proto.AuthService.GetEvents": {},
|
||||
"proto.AuthService.GetGithubConnectors": {},
|
||||
"proto.AuthService.GetInstallers": {},
|
||||
|
||||
@@ -63,6 +63,7 @@ import (
|
||||
dtauthz "github.com/gravitational/teleport/lib/devicetrust/authz"
|
||||
dtconfig "github.com/gravitational/teleport/lib/devicetrust/config"
|
||||
"github.com/gravitational/teleport/lib/events"
|
||||
iterstream "github.com/gravitational/teleport/lib/itertools/stream"
|
||||
"github.com/gravitational/teleport/lib/modules"
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
"github.com/gravitational/teleport/lib/services/local"
|
||||
@@ -6526,6 +6527,7 @@ func (a *ServerWithRoles) GetDatabase(ctx context.Context, name string) (types.D
|
||||
}
|
||||
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
func (a *ServerWithRoles) GetDatabases(ctx context.Context) (result []types.Database, err error) {
|
||||
if err := a.authorizeAction(types.KindDatabase, types.VerbList, types.VerbRead); err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
@@ -6543,6 +6545,42 @@ func (a *ServerWithRoles) GetDatabases(ctx context.Context) (result []types.Data
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
func (a *ServerWithRoles) ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error) {
|
||||
if err := a.authorizeAction(types.KindDatabase, types.VerbList, types.VerbRead); err != nil {
|
||||
return nil, "", trace.Wrap(err)
|
||||
}
|
||||
|
||||
if limit <= 0 || limit > apidefaults.DefaultChunkSize {
|
||||
limit = apidefaults.DefaultChunkSize
|
||||
}
|
||||
|
||||
var next string
|
||||
var seen int
|
||||
out, err := iterstream.Collect(
|
||||
iterstream.TakeWhile(
|
||||
iterstream.FilterMap(
|
||||
a.authServer.RangeDatabases(ctx, startKey, ""),
|
||||
func(db types.Database) (types.Database, bool) {
|
||||
if a.checkAccessToDatabase(db) == nil {
|
||||
return db, true
|
||||
}
|
||||
return nil, false
|
||||
},
|
||||
),
|
||||
func(db types.Database) bool {
|
||||
if seen < limit {
|
||||
seen++
|
||||
return true
|
||||
}
|
||||
next = db.GetName()
|
||||
return false
|
||||
},
|
||||
),
|
||||
)
|
||||
return out, next, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// DeleteDatabase removes the specified database resource.
|
||||
func (a *ServerWithRoles) DeleteDatabase(ctx context.Context, name string) error {
|
||||
if err := a.authorizeAction(types.KindDatabase, types.VerbDelete); err != nil {
|
||||
|
||||
@@ -78,6 +78,7 @@ import (
|
||||
"github.com/gravitational/teleport/lib/events"
|
||||
"github.com/gravitational/teleport/lib/events/eventstest"
|
||||
"github.com/gravitational/teleport/lib/integrations/awsra/createsession"
|
||||
"github.com/gravitational/teleport/lib/itertools/stream"
|
||||
"github.com/gravitational/teleport/lib/modules"
|
||||
"github.com/gravitational/teleport/lib/modules/modulestest"
|
||||
"github.com/gravitational/teleport/lib/okta/oktatest"
|
||||
@@ -2864,6 +2865,13 @@ func TestDatabasesCRUDRBAC(t *testing.T) {
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
dbs, next, err := devClt.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, next)
|
||||
require.Empty(t, cmp.Diff([]types.Database{devDatabase}, dbs,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
// Admin should see both.
|
||||
dbs, err = adminClt.GetDatabases(ctx)
|
||||
require.NoError(t, err)
|
||||
@@ -2871,6 +2879,21 @@ func TestDatabasesCRUDRBAC(t *testing.T) {
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
dbs, next, err = adminClt.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, next)
|
||||
require.Empty(t, cmp.Diff([]types.Database{adminDatabase, devDatabase}, dbs,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
// With limit, next should be dev
|
||||
dbs, next, err = adminClt.ListDatabases(ctx, 1, "")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, devDatabase.GetName(), next)
|
||||
require.Empty(t, cmp.Diff([]types.Database{adminDatabase}, dbs,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
// Dev shouldn't be able to delete dev database...
|
||||
err = devClt.DeleteDatabase(ctx, adminDatabase.GetName())
|
||||
require.True(t, trace.IsAccessDenied(err))
|
||||
@@ -2965,6 +2988,14 @@ func mustGetDatabases(t *testing.T, client *authclient.Client, wantDatabases []t
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
cmpopts.EquateEmpty(),
|
||||
))
|
||||
|
||||
actualDatabases, err = stream.Collect(client.RangeDatabases(t.Context(), "", ""))
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Empty(t, cmp.Diff(wantDatabases, actualDatabases,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
cmpopts.EquateEmpty(),
|
||||
))
|
||||
}
|
||||
|
||||
func TestKubernetesClusterCRUD_DiscoveryService(t *testing.T) {
|
||||
|
||||
@@ -277,8 +277,15 @@ type ReadProxyAccessPoint interface {
|
||||
GetDatabaseServers(ctx context.Context, namespace string, opts ...services.MarshalOption) ([]types.DatabaseServer, error)
|
||||
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
GetDatabases(ctx context.Context) ([]types.Database, error)
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error)
|
||||
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error]
|
||||
|
||||
// GetDatabase returns the specified database resource.
|
||||
GetDatabase(ctx context.Context, name string) (types.Database, error)
|
||||
|
||||
@@ -644,8 +651,15 @@ type ReadDatabaseAccessPoint interface {
|
||||
GetProxies() ([]types.Server, error)
|
||||
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
GetDatabases(ctx context.Context) ([]types.Database, error)
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error)
|
||||
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error]
|
||||
|
||||
// GetDatabase returns the specified database resource.
|
||||
GetDatabase(ctx context.Context, name string) (types.Database, error)
|
||||
}
|
||||
@@ -751,7 +765,15 @@ type ReadDiscoveryAccessPoint interface {
|
||||
GetKubernetesServers(ctx context.Context) ([]types.KubeServer, error)
|
||||
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
GetDatabases(ctx context.Context) ([]types.Database, error)
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error)
|
||||
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error]
|
||||
|
||||
// GetDatabase returns a database resource with the given name if it exists.
|
||||
GetDatabase(ctx context.Context, name string) (types.Database, error)
|
||||
|
||||
@@ -1100,8 +1122,15 @@ type Cache interface {
|
||||
GetDatabaseServers(ctx context.Context, namespace string, opts ...services.MarshalOption) ([]types.DatabaseServer, error)
|
||||
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
GetDatabases(ctx context.Context) ([]types.Database, error)
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error)
|
||||
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error]
|
||||
|
||||
// GetDatabase returns the specified database resource.
|
||||
GetDatabase(ctx context.Context, name string) (types.Database, error)
|
||||
|
||||
|
||||
@@ -4031,6 +4031,34 @@ func (g *GRPCServer) GetDatabases(ctx context.Context, _ *emptypb.Empty) (*types
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
func (g *GRPCServer) ListDatabases(ctx context.Context, req *authpb.ListDatabasesRequest) (*authpb.ListDatabasesResponse, error) {
|
||||
auth, err := g.authenticate(ctx)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
databases, next, err := auth.ListDatabases(ctx, int(req.PageSize), req.PageToken)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
|
||||
resp := &authpb.ListDatabasesResponse{
|
||||
Databases: make([]*types.DatabaseV3, 0, len(databases)),
|
||||
NextPageToken: next,
|
||||
}
|
||||
|
||||
for _, database := range databases {
|
||||
databaseV3, ok := database.(*types.DatabaseV3)
|
||||
if !ok {
|
||||
return nil, trace.BadParameter("unsupported database type %T", database)
|
||||
}
|
||||
resp.Databases = append(resp.Databases, databaseV3)
|
||||
}
|
||||
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
// DeleteDatabase removes the specified database.
|
||||
func (g *GRPCServer) DeleteDatabase(ctx context.Context, req *types.ResourceRequest) (*emptypb.Empty, error) {
|
||||
auth, err := g.authenticate(ctx)
|
||||
|
||||
@@ -3812,6 +3812,11 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
|
||||
out, next, err := clt.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
require.Empty(t, next)
|
||||
|
||||
// Create both databases.
|
||||
err = clt.CreateDatabase(ctx, db1)
|
||||
require.NoError(t, err)
|
||||
@@ -3825,6 +3830,13 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, next, err = clt.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, next)
|
||||
require.Empty(t, cmp.Diff([]types.Database{db1, db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
// Fetch a specific database.
|
||||
db, err := clt.GetDatabase(ctx, db2.GetName())
|
||||
require.NoError(t, err)
|
||||
@@ -3858,6 +3870,12 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
require.Empty(t, cmp.Diff([]types.Database{db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
out, next, err = clt.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, next)
|
||||
require.Empty(t, cmp.Diff([]types.Database{db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
// Try to delete a database that doesn't exist.
|
||||
err = clt.DeleteDatabase(ctx, "doesnotexist")
|
||||
@@ -3869,6 +3887,10 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
out, err = clt.GetDatabases(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
out, next, err = clt.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, next)
|
||||
require.Empty(t, out)
|
||||
}
|
||||
|
||||
// TestDatabaseServicesCRUD tests DatabaseService resource operations.
|
||||
|
||||
Vendored
+82
-3
@@ -18,6 +18,7 @@ package cache
|
||||
|
||||
import (
|
||||
"context"
|
||||
"iter"
|
||||
|
||||
"github.com/gravitational/trace"
|
||||
"google.golang.org/protobuf/proto"
|
||||
@@ -37,8 +38,8 @@ type databaseIndex string
|
||||
|
||||
const databaseNameIndex = "name"
|
||||
|
||||
func newDatabaseCollection(p services.Databases, w types.WatchKind) (*collection[types.Database, databaseIndex], error) {
|
||||
if p == nil {
|
||||
func newDatabaseCollection(upstream services.Databases, w types.WatchKind) (*collection[types.Database, databaseIndex], error) {
|
||||
if upstream == nil {
|
||||
return nil, trace.BadParameter("missing parameter Databases")
|
||||
}
|
||||
|
||||
@@ -52,7 +53,13 @@ func newDatabaseCollection(p services.Databases, w types.WatchKind) (*collection
|
||||
databaseNameIndex: types.Database.GetName,
|
||||
}),
|
||||
fetcher: func(ctx context.Context, loadSecrets bool) ([]types.Database, error) {
|
||||
return p.GetDatabases(ctx)
|
||||
out, err := stream.Collect(upstream.RangeDatabases(ctx, "", ""))
|
||||
// TODO(lokraszewski): DELETE IN v21.0.0
|
||||
if trace.IsNotImplemented(err) {
|
||||
out, err := upstream.GetDatabases(ctx)
|
||||
return out, trace.Wrap(err)
|
||||
}
|
||||
return out, trace.Wrap(err)
|
||||
},
|
||||
headerTransform: func(hdr *types.ResourceHeader) types.Database {
|
||||
return &types.DatabaseV3{
|
||||
@@ -92,6 +99,7 @@ func (c *Cache) GetDatabase(ctx context.Context, name string) (types.Database, e
|
||||
}
|
||||
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
func (c *Cache) GetDatabases(ctx context.Context) ([]types.Database, error) {
|
||||
ctx, span := c.Tracer.Start(ctx, "cache/GetDatabases")
|
||||
defer span.End()
|
||||
@@ -115,6 +123,77 @@ func (c *Cache) GetDatabases(ctx context.Context) ([]types.Database, error) {
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
func (c *Cache) ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error) {
|
||||
ctx, span := c.Tracer.Start(ctx, "cache/ListDatabases")
|
||||
defer span.End()
|
||||
|
||||
lister := genericLister[types.Database, databaseIndex]{
|
||||
cache: c,
|
||||
collection: c.collections.dbs,
|
||||
index: databaseNameIndex,
|
||||
upstreamList: c.Config.Databases.ListDatabases,
|
||||
nextToken: types.Database.GetName,
|
||||
}
|
||||
out, next, err := lister.list(ctx, limit, startKey)
|
||||
return out, next, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
func (c *Cache) RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error] {
|
||||
return func(yield func(types.Database, error) bool) {
|
||||
ctx, span := c.Tracer.Start(ctx, "cache/RangeDatabases")
|
||||
defer span.End()
|
||||
|
||||
rg, err := acquireReadGuard(c, c.collections.dbs)
|
||||
if err != nil {
|
||||
yield(nil, err)
|
||||
return
|
||||
}
|
||||
defer rg.Release()
|
||||
|
||||
if rg.ReadCache() {
|
||||
for database := range rg.store.resources(databaseNameIndex, start, end) {
|
||||
if !yield(database, nil) {
|
||||
return
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
rg.Release()
|
||||
|
||||
for database, err := range c.Config.Databases.RangeDatabases(ctx, start, end) {
|
||||
if err != nil {
|
||||
// TODO(lokraszewski): DELETE IN v21.0.0
|
||||
if trace.IsNotImplemented(err) {
|
||||
databases, err := c.Config.Databases.GetDatabases(ctx)
|
||||
if err != nil {
|
||||
yield(nil, err)
|
||||
return
|
||||
}
|
||||
|
||||
for _, database := range databases {
|
||||
if !yield(database, nil) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
yield(nil, err)
|
||||
return
|
||||
}
|
||||
|
||||
if !yield(database, nil) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
type databaseServerIndex string
|
||||
|
||||
const databaseServerNameIndex databaseServerIndex = "name"
|
||||
|
||||
Vendored
+110
-4
@@ -17,16 +17,25 @@
|
||||
package cache
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"context"
|
||||
"slices"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
gocmp "github.com/google/go-cmp/cmp"
|
||||
"github.com/google/go-cmp/cmp/cmpopts"
|
||||
"github.com/google/uuid"
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
apidefaults "github.com/gravitational/teleport/api/defaults"
|
||||
dbobjectv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/dbobject/v1"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/defaults"
|
||||
"github.com/gravitational/teleport/lib/itertools/stream"
|
||||
)
|
||||
|
||||
// TestDatabaseServices tests that CRUD operations on DatabaseServices are
|
||||
@@ -84,17 +93,114 @@ func TestDatabases(t *testing.T) {
|
||||
URI: "localhost:5432",
|
||||
})
|
||||
},
|
||||
create: p.databases.CreateDatabase,
|
||||
list: p.databases.GetDatabases,
|
||||
create: p.databases.CreateDatabase,
|
||||
list: func(ctx context.Context) ([]types.Database, error) {
|
||||
return stream.Collect(p.databases.RangeDatabases(ctx, "", ""))
|
||||
},
|
||||
cacheGet: p.cache.GetDatabase,
|
||||
cacheList: func(ctx context.Context, pageSize int) ([]types.Database, error) {
|
||||
return p.cache.GetDatabases(ctx)
|
||||
cacheList: func(ctx context.Context, _ int) ([]types.Database, error) {
|
||||
return stream.Collect(p.cache.RangeDatabases(ctx, "", ""))
|
||||
},
|
||||
update: p.databases.UpdateDatabase,
|
||||
deleteAll: p.databases.DeleteAllDatabases,
|
||||
})
|
||||
}
|
||||
|
||||
func TestDatabasesPagination(t *testing.T) {
|
||||
// TODO(okraport): extract this into generic helper for other paginated resources.
|
||||
t.Parallel()
|
||||
|
||||
p, err := newPack(t.TempDir(), ForProxy)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(p.Close)
|
||||
|
||||
ctx, cancel := context.WithCancel(t.Context())
|
||||
defer cancel()
|
||||
go func() {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-p.eventsC: // Drain events to prevent deadlocking.
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
expected := make([]types.Database, 0, 50)
|
||||
for i := range 50 {
|
||||
db, err := types.NewDatabaseV3(types.Metadata{
|
||||
Name: "db" + strconv.Itoa(i+1),
|
||||
}, types.DatabaseSpecV3{
|
||||
Protocol: defaults.ProtocolPostgres,
|
||||
URI: "localhost:5432",
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, p.databases.CreateDatabase(t.Context(), db))
|
||||
expected = append(expected, db)
|
||||
}
|
||||
slices.SortFunc(expected, func(a, b types.Database) int {
|
||||
return cmp.Compare(a.GetName(), b.GetName())
|
||||
})
|
||||
|
||||
// Wait for all the Databases to be replicated to the cache.
|
||||
require.EventuallyWithT(t, func(t *assert.CollectT) {
|
||||
assert.Equal(t, len(expected), p.cache.collections.dbs.store.len())
|
||||
}, 15*time.Second, 100*time.Millisecond)
|
||||
|
||||
out, err := p.cache.GetDatabases(t.Context())
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(expected))
|
||||
assert.Empty(t, gocmp.Diff(expected, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
page1, page2Start, err := p.cache.ListDatabases(t.Context(), 10, "")
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, page1, 10)
|
||||
assert.NotEmpty(t, page2Start)
|
||||
|
||||
page2, next, err := p.cache.ListDatabases(t.Context(), 1000, page2Start)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, page2, len(expected)-10)
|
||||
assert.Empty(t, next)
|
||||
|
||||
listed := append(page1, page2...)
|
||||
assert.Empty(t, gocmp.Diff(expected, listed,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, err = stream.Collect(p.cache.RangeDatabases(t.Context(), "", page2Start))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(page1))
|
||||
assert.Empty(t, gocmp.Diff(page1, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, err = stream.Collect(p.cache.RangeDatabases(t.Context(), "", ""))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(expected))
|
||||
assert.Empty(t, gocmp.Diff(expected, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, err = stream.Collect(p.cache.RangeDatabases(t.Context(), page2Start, ""))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(expected)-10)
|
||||
assert.Empty(t, gocmp.Diff(expected, append(page1, out...),
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
// invalidate the cache, cover upstream fallback
|
||||
p.cache.ok = false
|
||||
out, err = stream.Collect(p.cache.RangeDatabases(t.Context(), "", ""))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(expected))
|
||||
assert.Empty(t, gocmp.Diff(expected, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
}
|
||||
|
||||
// TestDatabaseServers tests that CRUD operations on database servers are
|
||||
// replicated from the backend to the cache.
|
||||
func TestDatabaseServers(t *testing.T) {
|
||||
|
||||
@@ -404,3 +404,24 @@ func MergeStreams[T any](
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TakeWhile iterates the stream taking items while predicate returns true
|
||||
func TakeWhile[T any](stream Stream[T], predicate func(T) bool) Stream[T] {
|
||||
return func(yield func(T, error) bool) {
|
||||
for item, err := range stream {
|
||||
if err != nil {
|
||||
yield(*new(T), trace.Wrap(err))
|
||||
return
|
||||
}
|
||||
|
||||
if !predicate(item) {
|
||||
return
|
||||
}
|
||||
|
||||
if !yield(item, nil) {
|
||||
return
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -884,3 +884,50 @@ func TestMergeStreams(t *testing.T) {
|
||||
require.Equal(t, []int{1, 2, 3, 4, 5, 6}, out)
|
||||
})
|
||||
}
|
||||
|
||||
func TestTakeWhile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Regular operation
|
||||
out, err := Collect(TakeWhile(
|
||||
Slice([]int{1, 2, 3, 4, 5, 6}),
|
||||
func(item int) bool {
|
||||
return item < 4
|
||||
},
|
||||
))
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []int{1, 2, 3}, out)
|
||||
|
||||
out, err = Collect(TakeWhile(
|
||||
Slice([]int{1, 2, 3, 4, 5, 6}),
|
||||
func(_ int) bool {
|
||||
return true
|
||||
},
|
||||
))
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []int{1, 2, 3, 4, 5, 6}, out)
|
||||
|
||||
// Propagate error
|
||||
out, err = Collect(TakeWhile(
|
||||
Fail[int](fmt.Errorf("unexpected error")),
|
||||
func(_ int) bool {
|
||||
return true
|
||||
},
|
||||
))
|
||||
|
||||
require.Error(t, err)
|
||||
require.Empty(t, out)
|
||||
|
||||
// Test early exit
|
||||
var actual []int
|
||||
TakeWhile(Slice([]int{1, 2, 3, 4, 5, 6}),
|
||||
func(item int) bool { return true },
|
||||
)(func(item int, err error) bool {
|
||||
actual = append(actual, item)
|
||||
return item < 3
|
||||
})
|
||||
require.Equal(t, []int{1, 2, 3}, actual)
|
||||
|
||||
}
|
||||
|
||||
@@ -21,6 +21,7 @@ package services
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"iter"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
@@ -46,7 +47,12 @@ import (
|
||||
// DatabaseGetter defines interface for fetching database resources.
|
||||
type DatabaseGetter interface {
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
GetDatabases(context.Context) ([]types.Database, error)
|
||||
// ListDatabases returns a page of database resources.
|
||||
ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error)
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error]
|
||||
// GetDatabase returns the specified database resource.
|
||||
GetDatabase(ctx context.Context, name string) (types.Database, error)
|
||||
}
|
||||
|
||||
@@ -20,25 +20,35 @@ package local
|
||||
|
||||
import (
|
||||
"context"
|
||||
"iter"
|
||||
"log/slog"
|
||||
|
||||
"github.com/gravitational/trace"
|
||||
|
||||
"github.com/gravitational/teleport"
|
||||
"github.com/gravitational/teleport/api/defaults"
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/backend"
|
||||
"github.com/gravitational/teleport/lib/itertools/stream"
|
||||
"github.com/gravitational/teleport/lib/services"
|
||||
)
|
||||
|
||||
// DatabaseService manages database resources in the backend.
|
||||
type DatabaseService struct {
|
||||
backend.Backend
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewDatabasesService creates a new DatabasesService.
|
||||
func NewDatabasesService(backend backend.Backend) *DatabaseService {
|
||||
return &DatabaseService{Backend: backend}
|
||||
return &DatabaseService{
|
||||
Backend: backend,
|
||||
logger: slog.With(teleport.ComponentKey, "DatabaseService"),
|
||||
}
|
||||
}
|
||||
|
||||
// GetDatabases returns all database resources.
|
||||
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
|
||||
func (s *DatabaseService) GetDatabases(ctx context.Context) ([]types.Database, error) {
|
||||
startKey := backend.ExactKey(databasesPrefix)
|
||||
result, err := s.GetRange(ctx, startKey, backend.RangeEnd(startKey), backend.NoLimit)
|
||||
@@ -57,6 +67,69 @@ func (s *DatabaseService) GetDatabases(ctx context.Context) ([]types.Database, e
|
||||
return databases, nil
|
||||
}
|
||||
|
||||
// ListDatabases returns a page of database resources.
|
||||
func (s *DatabaseService) ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error) {
|
||||
// Adjust page size, so it can't be too large.
|
||||
if limit <= 0 || limit > defaults.DefaultChunkSize {
|
||||
limit = defaults.DefaultChunkSize
|
||||
}
|
||||
|
||||
var next string
|
||||
var seen int
|
||||
out, err := stream.Collect(
|
||||
stream.TakeWhile(
|
||||
s.RangeDatabases(ctx, startKey, ""),
|
||||
func(db types.Database) bool {
|
||||
if seen < limit {
|
||||
seen++
|
||||
return true
|
||||
}
|
||||
next = db.GetName()
|
||||
return false
|
||||
},
|
||||
),
|
||||
)
|
||||
return out, next, trace.Wrap(err)
|
||||
}
|
||||
|
||||
// RangeDatabases returns database resources within the range [start, end).
|
||||
func (s *DatabaseService) RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error] {
|
||||
mapFn := func(item backend.Item) (types.Database, bool) {
|
||||
database, err := services.UnmarshalDatabase(item.Value,
|
||||
services.WithExpires(item.Expires),
|
||||
services.WithRevision(item.Revision))
|
||||
if err != nil {
|
||||
s.logger.WarnContext(ctx, "Failed to unmarshal database",
|
||||
"key", item.Key,
|
||||
"error", err,
|
||||
)
|
||||
return nil, false
|
||||
}
|
||||
return database, true
|
||||
}
|
||||
|
||||
dbKey := backend.NewKey(databasesPrefix)
|
||||
startKey := dbKey.AppendKey(backend.KeyFromString(start))
|
||||
endKey := backend.RangeEnd(dbKey)
|
||||
if end != "" {
|
||||
endKey = dbKey.AppendKey(backend.KeyFromString(end)).ExactKey()
|
||||
}
|
||||
|
||||
return stream.TakeWhile(
|
||||
stream.FilterMap(
|
||||
s.Backend.Items(ctx, backend.ItemsParams{
|
||||
StartKey: startKey,
|
||||
EndKey: endKey,
|
||||
}),
|
||||
mapFn,
|
||||
),
|
||||
func(db types.Database) bool {
|
||||
// The range is not inclusive of the end key, so return early
|
||||
// if the end has been reached.
|
||||
return end == "" || db.GetName() < end
|
||||
})
|
||||
}
|
||||
|
||||
// GetDatabase returns the specified database resource.
|
||||
func (s *DatabaseService) GetDatabase(ctx context.Context, name string) (types.Database, error) {
|
||||
item, err := s.Get(ctx, backend.NewKey(databasesPrefix, name))
|
||||
|
||||
@@ -19,18 +19,23 @@
|
||||
package local
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"context"
|
||||
"slices"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
gocmp "github.com/google/go-cmp/cmp"
|
||||
"github.com/google/go-cmp/cmp/cmpopts"
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/jonboulle/clockwork"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/gravitational/teleport/api/types"
|
||||
"github.com/gravitational/teleport/lib/backend/memory"
|
||||
"github.com/gravitational/teleport/lib/defaults"
|
||||
"github.com/gravitational/teleport/lib/itertools/stream"
|
||||
)
|
||||
|
||||
// TestDatabasesCRUD tests backend operations with database resources.
|
||||
@@ -66,6 +71,15 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
|
||||
out, next, err := service.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
require.Empty(t, next)
|
||||
|
||||
out, err = stream.Collect(service.RangeDatabases(ctx, "", ""))
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
|
||||
// Create both databases.
|
||||
err = service.CreateDatabase(ctx, db1)
|
||||
require.NoError(t, err)
|
||||
@@ -85,14 +99,27 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
// Fetch all databases.
|
||||
out, err = service.GetDatabases(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cmp.Diff([]types.Database{dbBadURI, db1, db2}, out,
|
||||
require.Empty(t, gocmp.Diff([]types.Database{dbBadURI, db1, db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, next, err = service.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, gocmp.Diff([]types.Database{dbBadURI, db1, db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
require.Empty(t, next)
|
||||
|
||||
out, err = stream.Collect(service.RangeDatabases(ctx, "", ""))
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, gocmp.Diff([]types.Database{dbBadURI, db1, db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
// Fetch a specific database.
|
||||
db, err := service.GetDatabase(ctx, db2.GetName())
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cmp.Diff(db2, db,
|
||||
require.Empty(t, gocmp.Diff(db2, db,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
@@ -110,7 +137,7 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
db, err = service.GetDatabase(ctx, db1.GetName())
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cmp.Diff(db1, db,
|
||||
require.Empty(t, gocmp.Diff(db1, db,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
@@ -119,7 +146,20 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
out, err = service.GetDatabases(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cmp.Diff([]types.Database{dbBadURI, db2}, out,
|
||||
require.Empty(t, gocmp.Diff([]types.Database{dbBadURI, db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, next, err = service.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, gocmp.Diff([]types.Database{dbBadURI, db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
require.Empty(t, next)
|
||||
|
||||
out, err = stream.Collect(service.RangeDatabases(ctx, "", ""))
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, gocmp.Diff([]types.Database{dbBadURI, db2}, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
@@ -133,4 +173,73 @@ func TestDatabasesCRUD(t *testing.T) {
|
||||
out, err = service.GetDatabases(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
|
||||
out, next, err = service.ListDatabases(ctx, 0, "")
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
require.Empty(t, next)
|
||||
|
||||
out, err = stream.Collect(service.RangeDatabases(ctx, "", ""))
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, out)
|
||||
|
||||
// Test pagination
|
||||
expected := make([]types.Database, 0, 50)
|
||||
for i := range 50 {
|
||||
db, err := types.NewDatabaseV3(types.Metadata{
|
||||
Name: "db" + strconv.Itoa(i+1),
|
||||
}, types.DatabaseSpecV3{
|
||||
Protocol: defaults.ProtocolPostgres,
|
||||
URI: "localhost",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, service.CreateDatabase(t.Context(), db))
|
||||
expected = append(expected, db)
|
||||
}
|
||||
slices.SortFunc(expected, func(a, b types.Database) int {
|
||||
return cmp.Compare(a.GetMetadata().Name, b.GetMetadata().Name)
|
||||
})
|
||||
|
||||
out, err = service.GetDatabases(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(expected))
|
||||
assert.Empty(t, gocmp.Diff(expected, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
page1, page2Start, err := service.ListDatabases(t.Context(), 10, "")
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, page1, 10)
|
||||
assert.NotEmpty(t, page2Start)
|
||||
|
||||
page2, next, err := service.ListDatabases(t.Context(), 1000, page2Start)
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, page2, len(expected)-10)
|
||||
assert.Empty(t, next)
|
||||
|
||||
listed := append(page1, page2...)
|
||||
assert.Empty(t, gocmp.Diff(expected, listed,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, err = stream.Collect(service.RangeDatabases(t.Context(), "", page2Start))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(page1))
|
||||
assert.Empty(t, gocmp.Diff(page1, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, err = stream.Collect(service.RangeDatabases(t.Context(), "", ""))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(expected))
|
||||
assert.Empty(t, gocmp.Diff(expected, out,
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
|
||||
out, err = stream.Collect(service.RangeDatabases(t.Context(), page2Start, ""))
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, out, len(expected)-10)
|
||||
assert.Empty(t, gocmp.Diff(expected, append(page1, out...),
|
||||
cmpopts.IgnoreFields(types.Metadata{}, "Revision"),
|
||||
))
|
||||
}
|
||||
|
||||
@@ -22,6 +22,7 @@ package discovery
|
||||
|
||||
import (
|
||||
"context"
|
||||
"iter"
|
||||
"maps"
|
||||
"slices"
|
||||
"testing"
|
||||
@@ -47,6 +48,7 @@ import (
|
||||
"github.com/gravitational/teleport/lib/auth/authclient"
|
||||
"github.com/gravitational/teleport/lib/backend/memory"
|
||||
"github.com/gravitational/teleport/lib/cloud/mocks"
|
||||
"github.com/gravitational/teleport/lib/itertools/stream"
|
||||
"github.com/gravitational/teleport/lib/services/local"
|
||||
"github.com/gravitational/teleport/lib/utils/log/logtest"
|
||||
)
|
||||
@@ -358,6 +360,14 @@ func (m *mockAuthServer) GetDatabases(ctx context.Context) ([]types.Database, er
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *mockAuthServer) ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error) {
|
||||
return nil, "", nil
|
||||
}
|
||||
|
||||
func (m *mockAuthServer) RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error] {
|
||||
return stream.Empty[types.Database]()
|
||||
}
|
||||
|
||||
func (m *mockAuthServer) GetNodes(ctx context.Context, namespace string) ([]types.Server, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -2021,9 +2021,17 @@ func (rc *ResourceCommand) Delete(ctx context.Context, client *authclient.Client
|
||||
}
|
||||
fmt.Printf("application %q has been deleted\n", rc.ref.Name)
|
||||
case types.KindDatabase:
|
||||
databases, err := client.GetDatabases(ctx)
|
||||
databases, err := stream.Collect(client.RangeDatabases(ctx, "", ""))
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
// TODO(okraport) DELETE IN v21.0.0
|
||||
if trace.IsNotImplemented(err) {
|
||||
databases, err = client.GetDatabases(ctx)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
} else {
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
}
|
||||
resDesc := "database"
|
||||
databases = filterByNameOrDiscoveredName(databases, rc.ref.Name)
|
||||
@@ -2842,10 +2850,19 @@ func (rc *ResourceCommand) getCollection(ctx context.Context, client *authclient
|
||||
}
|
||||
return &appCollection{apps: []types.Application{app}}, nil
|
||||
case types.KindDatabase:
|
||||
databases, err := client.GetDatabases(ctx)
|
||||
databases, err := stream.Collect(client.RangeDatabases(ctx, "", ""))
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
// TODO(okraport) DELETE IN v21.0.0
|
||||
if trace.IsNotImplemented(err) {
|
||||
databases, err = client.GetDatabases(ctx)
|
||||
if err != nil {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
} else {
|
||||
return nil, trace.Wrap(err)
|
||||
}
|
||||
}
|
||||
|
||||
if rc.ref.Name == "" {
|
||||
return &databaseCollection{databases: databases}, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user