diff --git a/lib/auth/grpcserver.go b/lib/auth/grpcserver.go index b4423f49eaf..03d5bd04388 100644 --- a/lib/auth/grpcserver.go +++ b/lib/auth/grpcserver.go @@ -110,6 +110,7 @@ import ( "github.com/gravitational/teleport/lib/services/local" "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/srv/server/installer" + usagereporter "github.com/gravitational/teleport/lib/usagereporter/teleport" "github.com/gravitational/teleport/lib/utils" ) @@ -5189,6 +5190,9 @@ func NewGRPCServer(cfg GRPCServerConfig) (*GRPCServer, error) { Authorizer: cfg.Authorizer, Backend: cfg.AuthServer.Services, Cache: cfg.AuthServer.Cache, + // This must be a function because cfg.AuthServer.UsageReporter is changed after `NewGRPCServer` is called. + // It starts as a DiscardUsageReporter, but when running in Cloud, gets replaced by a real reporter. + UsageReporter: func() usagereporter.UsageReporter { return cfg.AuthServer.UsageReporter }, }) if err != nil { return nil, trace.Wrap(err) diff --git a/lib/auth/usertasks/usertasksv1/service.go b/lib/auth/usertasks/usertasksv1/service.go index 07f1dc009a1..d36e411f8c2 100644 --- a/lib/auth/usertasks/usertasksv1/service.go +++ b/lib/auth/usertasks/usertasksv1/service.go @@ -26,8 +26,10 @@ import ( usertasksv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/usertasks/v1" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/types/usertasks" "github.com/gravitational/teleport/lib/authz" "github.com/gravitational/teleport/lib/services" + usagereporter "github.com/gravitational/teleport/lib/usagereporter/teleport" ) // ServiceConfig holds configuration options for the UserTask gRPC service. @@ -40,6 +42,9 @@ type ServiceConfig struct { // Cache is the cache for storing UserTask. Cache Reader + + // UsageReporter is the reporter for sending usage without it be related to an API call. + UsageReporter func() usagereporter.UsageReporter } // CheckAndSetDefaults checks the ServiceConfig fields and returns an error if @@ -55,6 +60,9 @@ func (s *ServiceConfig) CheckAndSetDefaults() error { if s.Cache == nil { return trace.BadParameter("cache is required") } + if s.UsageReporter == nil { + return trace.BadParameter("usage reporter is required") + } return nil } @@ -70,9 +78,10 @@ type Reader interface { type Service struct { usertasksv1.UnimplementedUserTaskServiceServer - authorizer authz.Authorizer - backend services.UserTasks - cache Reader + authorizer authz.Authorizer + backend services.UserTasks + cache Reader + usageReporter func() usagereporter.UsageReporter } // NewService returns a new UserTask gRPC service. @@ -82,9 +91,10 @@ func NewService(cfg ServiceConfig) (*Service, error) { } return &Service{ - authorizer: cfg.Authorizer, - backend: cfg.Backend, - cache: cfg.Cache, + authorizer: cfg.Authorizer, + backend: cfg.Backend, + cache: cfg.Cache, + usageReporter: cfg.UsageReporter, }, nil } @@ -104,9 +114,23 @@ func (s *Service) CreateUserTask(ctx context.Context, req *usertasksv1.CreateUse return nil, trace.Wrap(err) } + s.usageReporter().AnonymizeAndSubmit(userTaskToUserTaskStateEvent(req.GetUserTask())) + return rsp, nil } +func userTaskToUserTaskStateEvent(ut *usertasksv1.UserTask) *usagereporter.UserTaskStateEvent { + ret := &usagereporter.UserTaskStateEvent{ + TaskType: ut.GetSpec().GetTaskType(), + IssueType: ut.GetSpec().GetTaskType(), + State: ut.GetSpec().GetState(), + } + if ut.GetSpec().GetTaskType() == usertasks.TaskTypeDiscoverEC2 { + ret.InstancesCount = int32(len(ut.GetSpec().GetDiscoverEc2().GetInstances())) + } + return ret +} + // ListUserTasks returns a list of user tasks. func (s *Service) ListUserTasks(ctx context.Context, req *usertasksv1.ListUserTasksRequest) (*usertasksv1.ListUserTasksResponse, error) { authCtx, err := s.authorizer.Authorize(ctx) @@ -182,11 +206,20 @@ func (s *Service) UpdateUserTask(ctx context.Context, req *usertasksv1.UpdateUse return nil, trace.Wrap(err) } + existingUserTask, err := s.backend.GetUserTask(ctx, req.GetUserTask().GetMetadata().GetName()) + if err != nil { + return nil, trace.Wrap(err) + } + rsp, err := s.backend.UpdateUserTask(ctx, req.UserTask) if err != nil { return nil, trace.Wrap(err) } + if existingUserTask.GetSpec().GetState() != req.GetUserTask().GetSpec().GetState() { + s.usageReporter().AnonymizeAndSubmit(userTaskToUserTaskStateEvent(req.GetUserTask())) + } + return rsp, nil } @@ -201,11 +234,29 @@ func (s *Service) UpsertUserTask(ctx context.Context, req *usertasksv1.UpsertUse return nil, trace.Wrap(err) } + var emitStateChangeEvent bool + + existingUserTask, err := s.backend.GetUserTask(ctx, req.GetUserTask().GetMetadata().GetName()) + switch { + case trace.IsNotFound(err): + emitStateChangeEvent = true + + case err != nil: + return nil, trace.Wrap(err) + + default: + emitStateChangeEvent = existingUserTask.GetSpec().GetState() != req.GetUserTask().GetSpec().GetState() + } + rsp, err := s.backend.UpsertUserTask(ctx, req.UserTask) if err != nil { return nil, trace.Wrap(err) } + if emitStateChangeEvent { + s.usageReporter().AnonymizeAndSubmit(userTaskToUserTaskStateEvent(req.GetUserTask())) + } + return rsp, nil } diff --git a/lib/auth/usertasks/usertasksv1/service_test.go b/lib/auth/usertasks/usertasksv1/service_test.go index 48313c80556..d2e014476e0 100644 --- a/lib/auth/usertasks/usertasksv1/service_test.go +++ b/lib/auth/usertasks/usertasksv1/service_test.go @@ -29,15 +29,18 @@ import ( usertasksv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/usertasks/v1" "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/api/types/usertasks" "github.com/gravitational/teleport/lib/authz" "github.com/gravitational/teleport/lib/backend/memory" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/services/local" + usagereporter "github.com/gravitational/teleport/lib/usagereporter/teleport" "github.com/gravitational/teleport/lib/utils" ) func TestServiceAccess(t *testing.T) { t.Parallel() + testReporter := &mockUsageReporter{} testCases := []struct { name string @@ -78,7 +81,7 @@ func TestServiceAccess(t *testing.T) { t.Run(tt.name, func(t *testing.T) { for _, verbs := range utils.Combinations(tt.allowedVerbs) { t.Run(fmt.Sprintf("verbs=%v", verbs), func(t *testing.T) { - service := newService(t, fakeChecker{allowedVerbs: verbs}) + service := newService(t, fakeChecker{allowedVerbs: verbs}, testReporter) err := callMethod(t, service, tt.name) // expect access denied except with full set of verbs. if len(verbs) == len(tt.allowedVerbs) { @@ -105,6 +108,53 @@ func TestServiceAccess(t *testing.T) { }) } +func TestUsageEvents(t *testing.T) { + rwVerbs := []string{types.VerbList, types.VerbCreate, types.VerbRead, types.VerbUpdate, types.VerbDelete} + testReporter := &mockUsageReporter{} + service := newService(t, fakeChecker{allowedVerbs: rwVerbs}, testReporter) + ctx := context.Background() + + ut1, err := usertasks.NewDiscoverEC2UserTask(&usertasksv1.UserTaskSpec{ + Integration: "my-integration", + TaskType: "discover-ec2", + IssueType: "ec2-ssm-invocation-failure", + State: "OPEN", + DiscoverEc2: &usertasksv1.DiscoverEC2{ + AccountId: "123456789012", + Region: "us-east-1", + Instances: map[string]*usertasksv1.DiscoverEC2Instance{ + "i-123": &usertasksv1.DiscoverEC2Instance{ + InstanceId: "i-123", + DiscoveryConfig: "dc01", + DiscoveryGroup: "dg01", + }, + }, + }, + }) + require.NoError(t, err) + + _, err = service.CreateUserTask(ctx, &usertasksv1.CreateUserTaskRequest{UserTask: ut1}) + require.NoError(t, err) + // Usage reporting happens when user task is created, so we expect to see an event. + require.Len(t, testReporter.emittedEvents, 1) + + ut1.Spec.DiscoverEc2.Instances["i-345"] = &usertasksv1.DiscoverEC2Instance{ + InstanceId: "i-345", + DiscoveryConfig: "dc01", + DiscoveryGroup: "dg01", + } + _, err = service.UpsertUserTask(ctx, &usertasksv1.UpsertUserTaskRequest{UserTask: ut1}) + require.NoError(t, err) + // State was not updated, so usage events must not increase. + require.Len(t, testReporter.emittedEvents, 1) + + ut1.Spec.State = "RESOLVED" + _, err = service.UpdateUserTask(ctx, &usertasksv1.UpdateUserTaskRequest{UserTask: ut1}) + require.NoError(t, err) + // State was updated, so usage events include this new usage report. + require.Len(t, testReporter.emittedEvents, 2) +} + // callMethod calls a method with given name in the UserTask service func callMethod(t *testing.T, service *Service, method string) error { for _, desc := range usertasksv1.UserTaskService_ServiceDesc.Methods { @@ -132,7 +182,7 @@ func (f fakeChecker) CheckAccessToRule(_ services.RuleContext, _ string, resourc return trace.AccessDenied("access denied to rule=%v/verb=%v", resource, verb) } -func newService(t *testing.T, checker services.AccessChecker) *Service { +func newService(t *testing.T, checker services.AccessChecker, usageReporter usagereporter.UsageReporter) *Service { t.Helper() b, err := memory.New(memory.Config{}) @@ -153,10 +203,23 @@ func newService(t *testing.T, checker services.AccessChecker) *Service { }) service, err := NewService(ServiceConfig{ - Authorizer: authorizer, - Backend: backendService, - Cache: backendService, + Authorizer: authorizer, + Backend: backendService, + Cache: backendService, + UsageReporter: func() usagereporter.UsageReporter { return usageReporter }, }) require.NoError(t, err) return service } + +type mockUsageReporter struct { + emittedEvents []*usagereporter.UserTaskStateEvent +} + +func (m *mockUsageReporter) AnonymizeAndSubmit(events ...usagereporter.Anonymizable) { + for _, e := range events { + if userTaskEvent, ok := e.(*usagereporter.UserTaskStateEvent); ok { + m.emittedEvents = append(m.emittedEvents, userTaskEvent) + } + } +}