[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:
Luke Okraszewski
2025-08-19 10:05:40 +00:00
committed by GitHub
parent a987c5af9c
commit f04493e3ed
19 changed files with 2263 additions and 1063 deletions
+54
View File
@@ -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})
File diff suppressed because it is too large Load Diff
+43
View File
@@ -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": {},
+38
View File
@@ -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 {
+31
View File
@@ -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) {
+29
View File
@@ -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)
+28
View File
@@ -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)
+22
View File
@@ -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.
+82 -3
View File
@@ -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"
+110 -4
View File
@@ -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) {
+21
View File
@@ -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
}
}
}
}
+47
View File
@@ -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)
}
+6
View File
@@ -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)
}
+74 -1
View File
@@ -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))
+114 -5
View File
@@ -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"),
))
}
+10
View File
@@ -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
}
+21 -4
View File
@@ -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
}