feat: Add filtering and sort to tctl bots instances ls (#60273)

* Fix missing `--format` flag

* Use v2 rpc

* Add `--search` flag

* Add `--query` flag

* Add `--sort-index` and `--sort-order` flags

* Remove `generation` and add `version` fields to output

* Allow enabling the auth cache for the test process

* Add list bot instances tests

* Sync join method access logic between tctl and web

* Access `authentication.JoinMethod` safely

* Unhide `--format` flag

* Simplify version header label

* Fallback to v1 ListBotInstances

* Refactor to remove use of `authclient.ClientI`

* A way better fallback implementation 🙌

* typo 🙄

* Refactor to single interface
This commit is contained in:
Nick Marais
2025-10-20 11:48:16 +00:00
committed by GitHub
parent 8f8bcd438b
commit 5ceeda54c0
4 changed files with 377 additions and 47 deletions
+5 -1
View File
@@ -17,6 +17,7 @@
package web
import (
"cmp"
"context"
"fmt"
"net/http"
@@ -505,7 +506,10 @@ func (h *Handler) listBotInstancesV2(_ http.ResponseWriter, r *http.Request, _ h
}
if authentication != nil {
uiInstance.JoinMethodLatest = authentication.GetJoinMethod()
uiInstance.JoinMethodLatest = cmp.Or(
authentication.GetJoinAttrs().GetMeta().GetJoinMethod(),
authentication.GetJoinMethod(),
)
}
if heartbeat != nil {
+82 -44
View File
@@ -40,15 +40,17 @@ import (
"google.golang.org/protobuf/types/known/fieldmaskpb"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
machineidv1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1"
"github.com/gravitational/teleport/api/mfa"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/lib/asciitable"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/auth/machineid/machineidv1"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/set"
@@ -74,6 +76,12 @@ type BotsCommand struct {
addLogins string
setLogins string
search string
query string
sortIndex string
sortOrder string
botsList *kingpin.CmdClause
botsAdd *kingpin.CmdClause
botsRemove *kingpin.CmdClause
@@ -127,6 +135,11 @@ func (c *BotsCommand) Initialize(app *kingpin.Application, _ *tctlcfg.GlobalCLIF
c.botsInstancesList = c.botsInstances.Command("list", "List bot instances.").Alias("ls")
c.botsInstancesList.Arg("name", "The name of the bot from which to list instances. If unset, lists instances from all bots.").StringVar(&c.botName)
c.botsInstancesList.Flag("format", "Output format, 'text' or 'json'").Default(teleport.Text).EnumVar(&c.format, teleport.Text, teleport.JSON)
c.botsInstancesList.Flag("search", "Fuzzy search query used to filter bot instances").StringVar(&c.search)
c.botsInstancesList.Flag("query", "An expression in the Teleport predicate language used to filter bot instances").StringVar(&c.query)
c.botsInstancesList.Flag("sort-index", "Request sort index, 'bot_name', 'active_at_latest', 'version_latest' or 'host_name_latest'").Default("bot_name").StringVar(&c.sortIndex)
c.botsInstancesList.Flag("sort-order", "Request sort order, 'ascending' or 'descending'").Default("ascending").StringVar(&c.sortOrder)
c.botsInstancesAdd = c.botsInstances.Command("add", "Join a new instance onto an existing bot.").Alias("join")
c.botsInstancesAdd.Arg("name", "The name of the existing bot for which to add a new instance.").Required().StringVar(&c.botName)
@@ -139,8 +152,8 @@ func (c *BotsCommand) Initialize(app *kingpin.Application, _ *tctlcfg.GlobalCLIF
}
// TryRun attempts to run subcommands.
func (c *BotsCommand) TryRun(ctx context.Context, cmd string, clientFunc commonclient.InitFunc) (match bool, err error) {
var commandFunc func(ctx context.Context, client *authclient.Client) error
func (c *BotsCommand) TryRun(ctx context.Context, cmd string, clientFunc commonclient.InitFunc) (bool, error) {
var commandFunc func(ctx context.Context, client botsCommandClient) error
switch cmd {
case c.botsList.FullCommand():
commandFunc = c.ListBots
@@ -171,9 +184,22 @@ func (c *BotsCommand) TryRun(ctx context.Context, cmd string, clientFunc commonc
return true, trace.Wrap(err)
}
type botsCommandClient interface {
BotServiceClient() machineidv1pb.BotServiceClient
BotInstanceServiceClient() machineidv1pb.BotInstanceServiceClient
GetToken(ctx context.Context, name string) (types.ProvisionToken, error)
UpsertToken(ctx context.Context, token types.ProvisionToken) error
GetUser(ctx context.Context, name string, withSecrets bool) (types.User, error)
GetRole(context.Context, string) (types.Role, error)
UpsertLock(ctx context.Context, lock types.Lock) error
GetProxies() ([]types.Server, error)
PerformMFACeremony(ctx context.Context, in *proto.CreateAuthenticateChallengeRequest, promptOpts ...mfa.PromptOpt) (*proto.MFAAuthenticateResponse, error)
}
// ListBots writes a listing of the cluster's certificate renewal bots
// to standard out.
func (c *BotsCommand) ListBots(ctx context.Context, client *authclient.Client) error {
func (c *BotsCommand) ListBots(ctx context.Context, client botsCommandClient) error {
var bots []*machineidv1pb.Bot
req := &machineidv1pb.ListBotsRequest{}
for {
@@ -256,7 +282,7 @@ Please note:
`))
// AddBot adds a new certificate renewal bot to the cluster.
func (c *BotsCommand) AddBot(ctx context.Context, client *authclient.Client) error {
func (c *BotsCommand) AddBot(ctx context.Context, client botsCommandClient) error {
// Prompt for admin action MFA if required, allowing reuse for UpsertToken and CreateBot.
mfaResponse, err := mfa.PerformAdminActionMFACeremony(ctx, client.PerformMFACeremony, true /*allowReuse*/)
if err == nil {
@@ -345,7 +371,7 @@ func (c *BotsCommand) AddBot(ctx context.Context, client *authclient.Client) err
return trace.Wrap(outputToken(c.stdout, c.format, client, bot, token))
}
func (c *BotsCommand) RemoveBot(ctx context.Context, client *authclient.Client) error {
func (c *BotsCommand) RemoveBot(ctx context.Context, client botsCommandClient) error {
_, err := client.BotServiceClient().DeleteBot(ctx, &machineidv1pb.DeleteBotRequest{
BotName: c.botName,
})
@@ -358,7 +384,7 @@ func (c *BotsCommand) RemoveBot(ctx context.Context, client *authclient.Client)
return nil
}
func (c *BotsCommand) LockBot(ctx context.Context, client *authclient.Client) error {
func (c *BotsCommand) LockBot(ctx context.Context, client botsCommandClient) error {
lockExpiry, err := computeLockExpiry(c.lockExpires, c.lockTTL)
if err != nil {
return trace.Wrap(err)
@@ -453,14 +479,9 @@ func (c *BotsCommand) updateBotLogins(ctx context.Context, bot *machineidv1pb.Bo
return trace.Wrap(mask.Append(&machineidv1pb.Bot{}, "spec.traits"))
}
// clientRoleGetter is a minimal mockable interface for the client API
type clientRoleGetter interface {
GetRole(context.Context, string) (types.Role, error)
}
// updateBotRoles applies updates from CLI arguments to a bot's roles, updating
// the field mask as necessary if any updates were made.
func (c *BotsCommand) updateBotRoles(ctx context.Context, client clientRoleGetter, bot *machineidv1pb.Bot, mask *fieldmaskpb.FieldMask) error {
func (c *BotsCommand) updateBotRoles(ctx context.Context, client botsCommandClient, bot *machineidv1pb.Bot, mask *fieldmaskpb.FieldMask) error {
currentRoles := set.New[string](bot.Spec.Roles...)
var desiredRoles set.Set[string]
@@ -496,7 +517,7 @@ func (c *BotsCommand) updateBotRoles(ctx context.Context, client clientRoleGette
}
// UpdateBot performs various updates to existing bot users and roles.
func (c *BotsCommand) UpdateBot(ctx context.Context, client *authclient.Client) error {
func (c *BotsCommand) UpdateBot(ctx context.Context, client botsCommandClient) error {
bot, err := client.BotServiceClient().GetBot(ctx, &machineidv1pb.GetBotRequest{
BotName: c.botName,
})
@@ -547,27 +568,47 @@ func (c *BotsCommand) UpdateBot(ctx context.Context, client *authclient.Client)
}
// ListBotInstances lists bot instances, possibly filtering for a specific bot
func (c *BotsCommand) ListBotInstances(ctx context.Context, client *authclient.Client) error {
var instances []*machineidv1pb.BotInstance
req := &machineidv1pb.ListBotInstancesRequest{}
if c.botName != "" {
req.FilterBotName = c.botName
func (c *BotsCommand) ListBotInstances(ctx context.Context, client botsCommandClient) error {
pageFunc := func(ctx context.Context, pageSize int, pageToken string) ([]*machineidv1pb.BotInstance, string, error) {
resp, err := client.BotInstanceServiceClient().ListBotInstancesV2(ctx, &machineidv1pb.ListBotInstancesV2Request{
PageSize: int32(pageSize),
PageToken: pageToken,
SortField: c.sortIndex,
SortDesc: c.sortOrder == "descending",
Filter: &machineidv1pb.ListBotInstancesV2Request_Filters{
BotName: c.botName,
SearchTerm: c.search,
Query: c.query,
},
})
return resp.GetBotInstances(), resp.GetNextPageToken(), trace.Wrap(err)
}
for {
// TODO(nicholasmarais1158) Use ListBotInstancesV2 instead.
//nolint:staticcheck // SA1019
resp, err := client.BotInstanceServiceClient().ListBotInstances(ctx, req)
if err != nil {
return trace.Wrap(err)
fallbackFunc := func(ctx context.Context) ([]*machineidv1pb.BotInstance, error) {
if c.query != "" {
return nil, trace.NotImplemented("fallback not supported for requests with a query")
}
fallbackPageFunc := func(ctx context.Context, pageSize int, pageToken string) ([]*machineidv1pb.BotInstance, string, error) {
// Needed for backwards compatibility
//nolint:staticcheck // SA1019
resp, err := client.BotInstanceServiceClient().ListBotInstances(ctx, &machineidv1pb.ListBotInstancesRequest{
FilterBotName: c.botName,
PageSize: int32(pageSize),
PageToken: pageToken,
FilterSearchTerm: c.search,
Sort: &types.SortBy{
Field: c.sortIndex,
IsDesc: c.sortOrder == "descending",
},
})
return resp.GetBotInstances(), resp.GetNextPageToken(), trace.Wrap(err)
}
return stream.Collect(clientutils.Resources(ctx, fallbackPageFunc))
}
instances = append(instances, resp.BotInstances...)
if resp.NextPageToken == "" {
break
}
req.PageToken = resp.NextPageToken
instances, err := clientutils.CollectWithFallback(ctx, pageFunc, fallbackFunc)
if err != nil {
return trace.Wrap(err)
}
if c.format == teleport.JSON {
@@ -596,15 +637,14 @@ func (c *BotsCommand) ListBotInstances(ctx context.Context, client *authclient.C
return nil
}
t := asciitable.MakeTable([]string{"ID", "Join Method", "Hostname", "Joined", "Last Seen", "Generation"})
t := asciitable.MakeTable([]string{"ID", "Join Method", "Version", "Hostname", "Last Seen"})
for _, i := range instances {
var (
joinMethod string
hostname string
generation string
version string
)
joined := i.Status.InitialAuthentication.AuthenticatedAt.AsTime().Format(time.RFC3339)
initialJoinMethod := cmp.Or(
i.Status.InitialAuthentication.GetJoinAttrs().GetMeta().GetJoinMethod(),
i.Status.InitialAuthentication.JoinMethod,
@@ -612,13 +652,9 @@ func (c *BotsCommand) ListBotInstances(ctx context.Context, client *authclient.C
lastSeen := i.Status.InitialAuthentication.AuthenticatedAt.AsTime()
if len(i.Status.LatestAuthentications) == 0 {
generation = "n/a"
} else {
if len(i.Status.LatestAuthentications) > 0 {
auth := i.Status.LatestAuthentications[len(i.Status.LatestAuthentications)-1]
generation = fmt.Sprint(auth.Generation)
authJM := cmp.Or(
auth.GetJoinAttrs().GetMeta().GetJoinMethod(),
auth.JoinMethod,
@@ -636,11 +672,13 @@ func (c *BotsCommand) ListBotInstances(ctx context.Context, client *authclient.C
}
if len(i.Status.LatestHeartbeats) == 0 {
hostname = "n/a"
hostname = "-"
version = "-"
} else {
hb := i.Status.LatestHeartbeats[len(i.Status.LatestHeartbeats)-1]
hostname = hb.Hostname
version = hb.Version
if hb.RecordedAt.AsTime().After(lastSeen) {
lastSeen = hb.RecordedAt.AsTime()
@@ -649,7 +687,7 @@ func (c *BotsCommand) ListBotInstances(ctx context.Context, client *authclient.C
t.AddRow([]string{
fmt.Sprintf("%s/%s", i.Spec.BotName, i.Spec.InstanceId), joinMethod,
hostname, joined, lastSeen.Format(time.RFC3339), generation,
version, hostname, lastSeen.Format(time.RFC3339),
})
}
fmt.Fprintln(c.stdout, t.AsBuffer().String())
@@ -665,7 +703,7 @@ func (c *BotsCommand) ListBotInstances(ctx context.Context, client *authclient.C
}
// AddBotInstance begins onboarding a new instance of an existing bot.
func (c *BotsCommand) AddBotInstance(ctx context.Context, client *authclient.Client) error {
func (c *BotsCommand) AddBotInstance(ctx context.Context, client botsCommandClient) error {
// A bit of a misnomer but makes the terminology a bit more consistent. This
// doesn't directly create a bot instance, but creates token that allows a
// bot to join, which creates a new instance.
@@ -748,7 +786,7 @@ To onboard a new instance for this bot, run:
> {{.executable}} bots instances add {{.instance.Spec.BotName}}
`))
func (c *BotsCommand) ShowBotInstance(ctx context.Context, client *authclient.Client) error {
func (c *BotsCommand) ShowBotInstance(ctx context.Context, client botsCommandClient) error {
botName, instanceID, err := parseInstanceID(c.instanceID)
if err != nil {
return trace.Wrap(err)
@@ -801,7 +839,7 @@ type botJSONResponse struct {
}
// outputToken writes token information to stdout, depending on the token format.
func outputToken(wr io.Writer, format string, client *authclient.Client, bot *machineidv1pb.Bot, token types.ProvisionToken) error {
func outputToken(wr io.Writer, format string, client botsCommandClient, bot *machineidv1pb.Bot, token types.ProvisionToken) error {
if format == teleport.JSON {
tokenTTL := time.Duration(0)
if exp := token.Expiry(); !exp.IsZero() {
+282 -1
View File
@@ -24,10 +24,16 @@ import (
"slices"
"strings"
"testing"
"time"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"google.golang.org/protobuf/types/known/durationpb"
"google.golang.org/protobuf/types/known/fieldmaskpb"
"google.golang.org/protobuf/types/known/timestamppb"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/constants"
@@ -36,8 +42,11 @@ import (
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/integration/helpers"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/config"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/service"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/tool/teleport/testenv"
)
@@ -134,6 +143,7 @@ func TestUpdateBotLogins(t *testing.T) {
// mockAPIClient is a minimal API client used for testing
type mockRoleGetterClient struct {
*authclient.Client
roles []string
}
@@ -231,7 +241,7 @@ func TestUpdateBotRoles(t *testing.T) {
botRoles: tt.set,
}
err = cmd.updateBotRoles(context.TODO(), &mockClient, bot, fieldMask)
err = cmd.updateBotRoles(t.Context(), &mockClient, bot, fieldMask)
tt.assert(t, bot, fieldMask, err)
})
}
@@ -254,6 +264,7 @@ func TestAddAndListBotInstancesJSON(t *testing.T) {
ctx := context.Background()
client, err := testenv.NewDefaultAuthClient(process)
require.NoError(t, err)
t.Cleanup(func() { _ = client.Close() })
tokens, err := stream.Collect(clientutils.Resources(ctx, func(ctx context.Context, pageSize int, pageKey string) ([]types.ProvisionToken, string, error) {
@@ -304,3 +315,273 @@ func TestAddAndListBotInstancesJSON(t *testing.T) {
buf.Reset()
}
func TestListBotInstances(t *testing.T) {
t.Parallel()
dynAddr := helpers.NewDynamicServiceAddr(t)
fileConfig := &config.FileConfig{
Global: config.Global{
DataDir: t.TempDir(),
},
Auth: config.Auth{
Service: config.Service{
EnabledFlag: "true",
ListenAddress: dynAddr.AuthAddr,
},
},
}
process := makeAndRunTestAuthServer(t, withFileConfig(fileConfig), withFileDescriptors(dynAddr.Descriptors), withEnableCache(true))
ctx := t.Context()
client, err := testenv.NewDefaultAuthClient(process)
require.NoError(t, err)
t.Cleanup(func() { _ = client.Close() })
instance0 := createBotInstance(t, ctx, process)
instance1 := createBotInstance(t, ctx, process, func(instance *machineidv1pb.BotInstance) {
instance.Status.InitialHeartbeat.Hostname = "test-hostname-3"
instance.Status.InitialHeartbeat.Version = "19.0.1"
})
instance2 := createBotInstance(t, ctx, process, func(instance *machineidv1pb.BotInstance) {
instance.Spec.BotName = "test-bot-2"
instance.Status.InitialHeartbeat.Hostname = "test-hostname-2"
instance.Status.InitialHeartbeat.Version = "18.1.0"
})
// Give the auth cache a chance to catch-up
require.EventuallyWithT(t, func(t *assert.CollectT) {
res, _, err := process.GetAuthServer().ListBotInstances(ctx, 0, "", nil)
require.NoError(t, err)
require.Len(t, res, 3)
}, time.Second*10, time.Millisecond*50)
t.Run("defaults", func(t *testing.T) {
buf := strings.Builder{}
cmd := BotsCommand{
stdout: &buf,
format: teleport.JSON,
}
require.NoError(t, cmd.ListBotInstances(ctx, client))
res, err := services.UnmarshalProtoResourceArray[*machineidv1pb.BotInstance]([]byte(buf.String()))
require.NoError(t, err)
require.Len(t, res, 3)
})
t.Run("filter by bot name", func(t *testing.T) {
buf := strings.Builder{}
cmd := BotsCommand{
stdout: &buf,
format: teleport.JSON,
botName: "test-bot-1",
}
require.NoError(t, cmd.ListBotInstances(ctx, client))
res, err := services.UnmarshalProtoResourceArray[*machineidv1pb.BotInstance]([]byte(buf.String()))
require.NoError(t, err)
require.Len(t, res, 2)
assertContainsInstance(t, res, instance0.GetSpec().GetInstanceId())
assertContainsInstance(t, res, instance1.GetSpec().GetInstanceId())
})
t.Run("filter with search", func(t *testing.T) {
buf := strings.Builder{}
cmd := BotsCommand{
stdout: &buf,
format: teleport.JSON,
search: "test-hostname-2",
}
require.NoError(t, cmd.ListBotInstances(ctx, client))
res, err := services.UnmarshalProtoResourceArray[*machineidv1pb.BotInstance]([]byte(buf.String()))
require.NoError(t, err)
require.Len(t, res, 1)
assertContainsInstance(t, res, instance2.GetSpec().GetInstanceId())
})
t.Run("filter with query", func(t *testing.T) {
buf := strings.Builder{}
cmd := BotsCommand{
stdout: &buf,
format: teleport.JSON,
query: `status.latest_heartbeat.hostname == "test-hostname-2"`,
}
require.NoError(t, cmd.ListBotInstances(ctx, client))
res, err := services.UnmarshalProtoResourceArray[*machineidv1pb.BotInstance]([]byte(buf.String()))
require.NoError(t, err)
require.Len(t, res, 1)
assertContainsInstance(t, res, instance2.GetSpec().GetInstanceId())
})
t.Run("sort by field", func(t *testing.T) {
buf := strings.Builder{}
cmd := BotsCommand{
stdout: &buf,
format: teleport.JSON,
sortIndex: "version_latest",
}
require.NoError(t, cmd.ListBotInstances(ctx, client))
res, err := services.UnmarshalProtoResourceArray[*machineidv1pb.BotInstance]([]byte(buf.String()))
require.NoError(t, err)
require.Len(t, res, 3)
assert.Equal(t, "18.1.0", res[0].GetStatus().GetInitialHeartbeat().GetVersion())
assert.Equal(t, "19.0.0", res[1].GetStatus().GetInitialHeartbeat().GetVersion())
assert.Equal(t, "19.0.1", res[2].GetStatus().GetInitialHeartbeat().GetVersion())
})
t.Run("sort order", func(t *testing.T) {
buf := strings.Builder{}
cmd := BotsCommand{
stdout: &buf,
format: teleport.JSON,
sortIndex: "version_latest",
sortOrder: "descending",
}
require.NoError(t, cmd.ListBotInstances(ctx, client))
res, err := services.UnmarshalProtoResourceArray[*machineidv1pb.BotInstance]([]byte(buf.String()))
require.NoError(t, err)
require.Len(t, res, 3)
assert.Equal(t, "19.0.1", res[0].GetStatus().GetInitialHeartbeat().GetVersion())
assert.Equal(t, "19.0.0", res[1].GetStatus().GetInitialHeartbeat().GetVersion())
assert.Equal(t, "18.1.0", res[2].GetStatus().GetInitialHeartbeat().GetVersion())
})
}
func assertContainsInstance(t *testing.T, res []*machineidv1pb.BotInstance, instanceId string) {
assert.True(t, slices.ContainsFunc(res, func(in *machineidv1pb.BotInstance) bool {
return in.GetSpec().GetInstanceId() == instanceId
}))
}
func createBotInstance(t *testing.T, ctx context.Context, process *service.TeleportProcess, options ...func(instance *machineidv1pb.BotInstance)) (result *machineidv1pb.BotInstance) {
heartbeat := &machineidv1pb.BotInstanceStatusHeartbeat{
RecordedAt: timestamppb.New(time.Now()),
IsStartup: true,
Version: "19.0.0",
Hostname: "test-hostname-1",
Uptime: durationpb.New(1 * time.Hour),
Os: "linux",
}
base := &machineidv1pb.BotInstance{
Spec: &machineidv1pb.BotInstanceSpec{
BotName: "test-bot-1",
InstanceId: uuid.New().String(),
},
Status: &machineidv1pb.BotInstanceStatus{
InitialHeartbeat: heartbeat,
LatestHeartbeats: []*machineidv1pb.BotInstanceStatusHeartbeat{
heartbeat,
},
},
}
for _, fn := range options {
fn(base)
}
result, err := process.GetAuthServer().CreateBotInstance(ctx, base)
require.NoError(t, err)
return
}
func TestListBotInstancesFallback(t *testing.T) {
t.Parallel()
dynAddr := helpers.NewDynamicServiceAddr(t)
fileConfig := &config.FileConfig{
Global: config.Global{
DataDir: t.TempDir(),
},
Auth: config.Auth{
Service: config.Service{
EnabledFlag: "true",
ListenAddress: dynAddr.AuthAddr,
},
},
}
process := makeAndRunTestAuthServer(t, withFileConfig(fileConfig), withFileDescriptors(dynAddr.Descriptors), withEnableCache(true))
ctx := t.Context()
client, err := testenv.NewDefaultAuthClient(process)
require.NoError(t, err)
authClient := &mockBotInstanceListerClient{
Client: client,
}
t.Run("fallback allowed", func(t *testing.T) {
cmd := BotsCommand{
stdout: ptr(strings.Builder{}),
format: teleport.JSON,
}
require.NoError(t, cmd.ListBotInstances(ctx, authClient))
})
t.Run("fallback not allowed", func(t *testing.T) {
cmd := BotsCommand{
stdout: ptr(strings.Builder{}),
format: teleport.JSON,
query: "foo()", // query is only available in ListBotInstancesV2
}
err := cmd.ListBotInstances(ctx, authClient)
require.Error(t, err)
require.ErrorContains(t, err, "fallback not supported for requests with a query")
})
}
// mockBotInstanceListerClient is a client which returns NotImplemented for
// ListBotInstancesV2 to simulate a service running an older version.
type mockBotInstanceListerClient struct {
*authclient.Client
}
func (c *mockBotInstanceListerClient) BotInstanceServiceClient() machineidv1pb.BotInstanceServiceClient {
return &mockBotInstanceListV2ErrorClient{
BotInstanceServiceClient: c.Client.BotInstanceServiceClient(),
errV1: nil,
errV2: trace.NotImplemented("not implemeted in mock"),
}
}
type mockBotInstanceListV2ErrorClient struct {
machineidv1pb.BotInstanceServiceClient
errV1 error
errV2 error
}
func (c *mockBotInstanceListV2ErrorClient) ListBotInstances(ctx context.Context, in *machineidv1pb.ListBotInstancesRequest, opts ...grpc.CallOption) (*machineidv1pb.ListBotInstancesResponse, error) {
if c.errV1 == nil {
// Needed for backwards compatibility
//nolint:staticcheck // SA1019
return c.BotInstanceServiceClient.ListBotInstances(ctx, in, opts...)
}
return nil, c.errV2
}
func (c *mockBotInstanceListV2ErrorClient) ListBotInstancesV2(ctx context.Context, in *machineidv1pb.ListBotInstancesV2Request, opts ...grpc.CallOption) (*machineidv1pb.ListBotInstancesResponse, error) {
if c.errV2 == nil {
return c.BotInstanceServiceClient.ListBotInstancesV2(ctx, in, opts...)
}
return nil, c.errV2
}
func ptr[T any](v T) *T { return &v }
+8 -1
View File
@@ -217,6 +217,7 @@ type testServerOptions struct {
fileConfig *config.FileConfig
fileDescriptors []*servicecfg.FileDescriptor
fakeClock *clockwork.FakeClock
enableCache bool
}
type testServerOptionFunc func(options *testServerOptions)
@@ -239,6 +240,12 @@ func withFakeClock(fakeClock *clockwork.FakeClock) testServerOptionFunc {
}
}
func withEnableCache(enableCache bool) testServerOptionFunc {
return func(options *testServerOptions) {
options.enableCache = enableCache
}
}
func makeAndRunTestAuthServer(t *testing.T, opts ...testServerOptionFunc) (auth *service.TeleportProcess) {
var options testServerOptions
for _, opt := range opts {
@@ -254,7 +261,7 @@ func makeAndRunTestAuthServer(t *testing.T, opts ...testServerOptionFunc) (auth
require.NoError(t, err)
}
cfg.CachePolicy.Enabled = false
cfg.CachePolicy.Enabled = options.enableCache
cfg.Proxy.DisableWebInterface = true
cfg.InstanceMetadataClient = imds.NewDisabledIMDSClient()
if options.fakeClock != nil {