diff --git a/integration/integration_test.go b/integration/integration_test.go index 8624b8a7da9..4b4e53285bf 100644 --- a/integration/integration_test.go +++ b/integration/integration_test.go @@ -446,17 +446,17 @@ func testAuditOn(t *testing.T, suite *integrationTestSuite) { require.NoError(t, err) require.Empty(t, sessions) + cl, err := teleport.NewClient(helpers.ClientConfig{ + Login: suite.Me.Username, + Cluster: helpers.Site, + Host: nodeConf.Hostname, + Port: helpers.Port(t, nodeConf.SSH.Addr.Addr), + ForwardAgent: tt.inForwardAgent, + }) // create interactive session (this goroutine is this user's terminal time) endC := make(chan error) myTerm := NewTerminal(250) go func() { - cl, err := teleport.NewClient(helpers.ClientConfig{ - Login: suite.Me.Username, - Cluster: helpers.Site, - Host: nodeConf.Hostname, - Port: helpers.Port(t, nodeConf.SSH.Addr.Addr), - ForwardAgent: tt.inForwardAgent, - }) if err != nil { endC <- err return @@ -503,8 +503,11 @@ func testAuditOn(t *testing.T, suite *integrationTestSuite) { } } + cc, err := cl.ConnectToCluster(ctx) + require.NoError(t, err) + t.Cleanup(func() { cc.Close() }) // Test streaming events and recording. - capturedStream, sessionEvents := streamSession(ctx, t, site, sessionID) + capturedStream, sessionEvents := streamSession(ctx, t, cc.AuthClient, sessionID) findByType := func(et string) apievents.AuditEvent { for _, e := range sessionEvents { @@ -542,7 +545,7 @@ func testAuditOn(t *testing.T, suite *integrationTestSuite) { require.Regexp(t, ".*exit.*", recorded) require.Regexp(t, ".*echo hi.*", recorded) - sessionEvents, _, err = site.SearchEvents(ctx, events.SearchEventsRequest{ + sessionEvents, _, err = cc.AuthClient.SearchEvents(ctx, events.SearchEventsRequest{ From: time.Time{}, To: time.Now(), EventTypes: []string{ @@ -2384,6 +2387,9 @@ func twoClustersTunnel(t *testing.T, suite *integrationTestSuite, now time.Time, // wait for active tunnel connections to be established helpers.WaitForActiveTunnelConnections(t, b.Tunnel, a.Secrets.SiteName, 1) + err = b.WaitForNodeCount(ctx, a.Secrets.SiteName, 1) + require.NoError(t, err) + // via tunnel b->a: tc, err = b.NewClient(helpers.ClientConfig{ Login: username, @@ -2402,12 +2408,12 @@ func twoClustersTunnel(t *testing.T, suite *integrationTestSuite, now time.Time, }, 10*time.Second, 250*time.Millisecond) require.Equal(t, "hello world\n", stdout.String()) - clientHasEvents := func(site authclient.ClientI, count int) func() bool { + clientHasEvents := func(cc authclient.ClientI, count int) func() bool { // only look for exec events eventTypes := []string{events.ExecEvent} return func() bool { - eventsInSite, _, err := site.SearchEvents(ctx, events.SearchEventsRequest{ + eventsInSite, _, err := cc.SearchEvents(ctx, events.SearchEventsRequest{ From: now, To: now.Add(1 * time.Hour), EventTypes: eventTypes, @@ -2418,10 +2424,20 @@ func twoClustersTunnel(t *testing.T, suite *integrationTestSuite, now time.Time, } } - siteA := a.GetSiteAPI(a.Secrets.SiteName) + tA, err := a.NewClient(helpers.ClientConfig{ + Login: username, + Cluster: a.Secrets.SiteName, + Host: Host, + Port: sshPort, + ForwardAgent: true, + }) + require.NoError(t, err) + cA, err := tA.ConnectToCluster(ctx) + require.NoError(t, err) + t.Cleanup(func() { cA.Close() }) // Wait for 2nd event before stopping auth. - require.Eventually(t, clientHasEvents(siteA, 2), 5*time.Second, 500*time.Millisecond, + require.Eventually(t, clientHasEvents(cA.AuthClient, 2), 5*time.Second, 500*time.Millisecond, "Failed to find %d events on helpers.Site A after 5s", execCountSiteA) // Stop "site-A" and try to connect to it again via "site-A" (expect a connection error) @@ -2445,12 +2461,21 @@ func twoClustersTunnel(t *testing.T, suite *integrationTestSuite, now time.Time, require.Eventually(t, tcHasReconnected, 10*time.Second, 250*time.Millisecond, "Timed out waiting for helpers.Site A to restart: %v", sshErr) - siteA = a.GetSiteAPI(a.Secrets.SiteName) - require.Eventually(t, clientHasEvents(siteA, execCountSiteA), 5*time.Second, 500*time.Millisecond, + require.Eventually(t, clientHasEvents(cA.AuthClient, execCountSiteA), 5*time.Second, 500*time.Millisecond, "Failed to find %d events on helpers.Site A after 5s", execCountSiteA) - siteB := b.GetSiteAPI(b.Secrets.SiteName) - require.Eventually(t, clientHasEvents(siteB, execCountSiteB), 5*time.Second, 500*time.Millisecond, + bClient, err := b.NewClient(helpers.ClientConfig{ + Login: username, + Cluster: b.Secrets.SiteName, + Host: Host, + Port: sshPort, + ForwardAgent: true, + }) + require.NoError(t, err) + cB, err := bClient.ConnectToCluster(ctx) + require.NoError(t, err) + t.Cleanup(func() { cB.Close() }) + require.Eventually(t, clientHasEvents(cB.AuthClient, execCountSiteB), 5*time.Second, 500*time.Millisecond, "Failed to find %d events on helpers.Site B after 5s", execCountSiteB) } @@ -4927,17 +4952,14 @@ func testAuditOff(t *testing.T, suite *integrationTestSuite) { endCh := make(chan error, 1) myTerm := NewTerminal(250) + cl, err := teleport.NewClient(helpers.ClientConfig{ + Login: suite.Me.Username, + Cluster: helpers.Site, + Host: Host, + Port: helpers.Port(t, teleport.SSH), + }) + require.NoError(t, err) go func() { - cl, err := teleport.NewClient(helpers.ClientConfig{ - Login: suite.Me.Username, - Cluster: helpers.Site, - Host: Host, - Port: helpers.Port(t, teleport.SSH), - }) - if err != nil { - endCh <- err - return - } cl.Stdout = myTerm cl.Stdin = myTerm err = cl.SSH(ctx, []string{}) @@ -4965,7 +4987,10 @@ func testAuditOff(t *testing.T, suite *integrationTestSuite) { // however, attempts to read the actual sessions should fail because it was // not actually recorded - eventsCh, errCh := site.StreamSessionEvents(ctx, session.ID(tracker.GetSessionID()), 0) + cc, err := cl.ConnectToCluster(ctx) + require.NoError(t, err) + t.Cleanup(func() { cc.Close() }) + eventsCh, errCh := cc.AuthClient.StreamSessionEvents(ctx, session.ID(tracker.GetSessionID()), 0) err = nil readLoop: for { @@ -4983,7 +5008,7 @@ readLoop: // ensure that session related events were emitted to audit log var auditEvents []apievents.AuditEvent require.Eventually(t, func() bool { - ae, _, err := site.SearchEvents(ctx, events.SearchEventsRequest{ + ae, _, err := cc.AuthClient.SearchEvents(ctx, events.SearchEventsRequest{ From: beforeSession, To: time.Now(), EventTypes: []string{ @@ -7349,6 +7374,16 @@ func testSessionStreaming(t *testing.T, suite *integrationTestSuite) { defer teleport.StopAll() api := teleport.GetSiteAPI(helpers.Site) + cl, err := teleport.NewClient(helpers.ClientConfig{ + Login: suite.Me.Username, + Cluster: helpers.Site, + Host: Host, + Port: helpers.Port(t, teleport.SSH), + }) + require.NoError(t, err) + clusterClient, err := cl.ConnectToCluster(ctx) + require.NoError(t, err) + t.Cleanup(func() { clusterClient.Close() }) uploadStream, err := api.CreateAuditStream(ctx, sessionID) require.NoError(t, err) @@ -7373,7 +7408,9 @@ outer: time.Sleep(time.Second * 5) receivedSession := make([]apievents.AuditEvent, 0) - sessionPlayback, e := api.StreamSessionEvents(ctx, sessionID, 0) + // StreamSessionEvents can no longer be called by builtin Teleport identities, so + // we need to stream using a ClusterClient + sessionPlayback, e := clusterClient.AuthClient.StreamSessionEvents(ctx, sessionID, 0) inner: for { diff --git a/lib/auth/auth.go b/lib/auth/auth.go index f73f5927a9a..a63e51041f7 100644 --- a/lib/auth/auth.go +++ b/lib/auth/auth.go @@ -147,7 +147,7 @@ import ( "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/services/local" "github.com/gravitational/teleport/lib/services/readonly" - "github.com/gravitational/teleport/lib/session" + libsession "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/sshca" "github.com/gravitational/teleport/lib/sshutils" "github.com/gravitational/teleport/lib/tlsca" @@ -2552,6 +2552,32 @@ func (a *Server) GetClock() clockwork.Clock { return a.clock } +// OnUploadComplete is called after a session recording upload completes. It +// streams the session events to find the existing session end event, or +// reconstructs and emits one from the session start event if none is found. +// It is the canonical OnUploadComplete callback used by the ProtoStreamer, +// AuditLog, and recording encryption service. +// +// TODO(tigrato): this check is not 100% correct. Instead of streaming the file, +// one should query the audit log to ensure the event exists. The file can contain +// the session end event but for some reason the the audit log event was lost. +// There are many reasons for that to happen such as audit queue in the agent being +// full, audit backend being down, agent restarting after writing the session complete. +func (a *Server) OnUploadComplete(ctx context.Context, sessionID libsession.ID) (apievents.AuditEvent, error) { + clusterName, err := a.GetClusterName(ctx) + if err != nil { + return nil, trace.Wrap(err) + } + return events.FindOrRecoverSessionEnd(ctx, events.FindOrRecoverSessionEndConfig{ + ClusterName: clusterName.GetClusterName(), + Streamer: a, + SessionID: sessionID, + AuditLog: a, + Log: a.logger, + Clock: a.clock, + }) +} + // SetBcryptCost sets bcryptCostOverride, used in tests func (a *Server) SetBcryptCost(cost int) { a.lock.Lock() @@ -8885,7 +8911,7 @@ func (s *Server) GetSigstorePolicyEvaluator() workloadidentityv1.SigstorePolicyE } // TODO(tigrato): remove Download* methods once e no longer references them. -func (s *Server) DownloadSummary(ctx context.Context, sessionID session.ID, writer io.Writer) error { +func (s *Server) DownloadSummary(ctx context.Context, sessionID libsession.ID, writer io.Writer) error { reader, err := s.StreamSessionSummary(ctx, sessionID) if err != nil { return trace.Wrap(err) diff --git a/lib/auth/auth_with_roles.go b/lib/auth/auth_with_roles.go index 26bdba8a484..c0ffb710c3d 100644 --- a/lib/auth/auth_with_roles.go +++ b/lib/auth/auth_with_roles.go @@ -259,6 +259,15 @@ func (a *ServerWithRoles) actionForKindSession(ctx context.Context, sid session. return nil } + // Fast pre-check: if no role even mentions KindSession/VerbRead, skip the + // more expensive predicate evaluation below. + if err := a.context.Checker.GuessIfAccessIsPossible( + &services.Context{User: a.context.User}, + apidefaults.Namespace, types.KindSession, types.VerbRead, + ); err != nil { + return trace.Wrap(err) + } + // First try a simple check without the extended context. if err := a.actionWithContext(&services.Context{User: a.context.User}, types.KindSession, types.VerbRead); err == nil { return nil @@ -6738,15 +6747,6 @@ func (a *ServerWithRoles) ReplaceRemoteLocks(ctx context.Context, clusterName st // channel if one is encountered. Otherwise the event channel is closed when the stream ends. // The event channel is not closed on error to prevent race conditions in downstream select statements. func (a *ServerWithRoles) StreamSessionEvents(ctx context.Context, sessionID session.ID, startIndex int64) (chan apievents.AuditEvent, chan error) { - err := a.localServerAction() - isTeleportServer := err == nil - - // StreamSessionEvents can be called internally, and when that - // happens we don't want to emit an event or check for permissions. - if isTeleportServer { - return a.alog.StreamSessionEvents(ctx, sessionID, startIndex) - } - if err := a.actionForKindSession(ctx, sessionID); err != nil { c, e := make(chan apievents.AuditEvent), make(chan error, 1) e <- trace.Wrap(err) diff --git a/lib/auth/auth_with_roles_test.go b/lib/auth/auth_with_roles_test.go index d8cf2c665ba..ad3f484d704 100644 --- a/lib/auth/auth_with_roles_test.go +++ b/lib/auth/auth_with_roles_test.go @@ -2630,34 +2630,33 @@ func TestStreamSessionEvents_User(t *testing.T) { require.Equal(t, username, event.User) } -// TestStreamSessionEvents_Builtin ensures that when a builtin role streams a session's events, it does not emit -// an audit event. -func TestStreamSessionEvents_Builtin(t *testing.T) { +// TestStreamSessionEvents_Builtin ensures that a builtin role can not stream a session's events +// or read audit log. +func TestAuditLog_SessionEvents_BuiltinRoles(t *testing.T) { t.Parallel() ctx := t.Context() srv := newTestTLSServer(t) - identity := authtest.TestBuiltin(types.RoleProxy) - clt, err := srv.NewClient(identity) - require.NoError(t, err) + roles := types.LocalServiceMappings() + for _, role := range roles { + identity := authtest.TestBuiltin(role) + clt, err := srv.NewClient(identity) + require.NoError(t, err) - // ignore the response as we don't want the events or the error (the session will not exist) - _, _ = clt.StreamSessionEvents(ctx, "44c6cea8-362f-11ea-83aa-125400432324", 0) + _, errCh := clt.StreamSessionEvents(ctx, "44c6cea8-362f-11ea-83aa-125400432324", 0) + require.True(t, trace.IsAccessDenied(<-errCh), "expected access denied error when streaming for builtin role") - // we need to wait for a short period to ensure the event is returned - time.Sleep(500 * time.Millisecond) - - searchEvents, _, err := srv.AuthServer.AuditLog.SearchEvents(ctx, events.SearchEventsRequest{ - From: srv.Clock().Now().Add(-time.Hour), - To: srv.Clock().Now().Add(time.Hour), - EventTypes: []string{events.SessionRecordingAccessEvent}, - Limit: 1, - Order: types.EventOrderDescending, - }) - require.NoError(t, err) - - require.Empty(t, searchEvents) + _, _, err = clt.SearchEvents(ctx, events.SearchEventsRequest{ + From: srv.Clock().Now().Add(-time.Hour), + To: srv.Clock().Now().Add(time.Hour), + EventTypes: []string{events.SessionRecordingAccessEvent}, + Limit: 1, + Order: types.EventOrderDescending, + }) + require.Error(t, err) + require.True(t, trace.IsAccessDenied(err), "expected access denied error when streaming for builtin role") + } } // TestStreamSessionEvents ensures that when a user streams a session's events @@ -2780,6 +2779,79 @@ func TestStreamSessionEvents_SessionType(t *testing.T) { require.Equal(t, accessedFormat, event.Format) } +// TestOnUploadComplete_RecoversMissingSessionEnd verifies that +// auth.Server.OnUploadComplete finds and emits a recovered session end event +// when the session recording was completed without one (e.g. due to a crash). +func TestOnUploadComplete_RecoversMissingSessionEnd(t *testing.T) { + t.Parallel() + ctx := t.Context() + + authServerConfig := authtest.AuthServerConfig{ + Dir: t.TempDir(), + Clock: clockwork.NewFakeClockAt(time.Now().Round(time.Second).UTC()), + } + require.NoError(t, authServerConfig.CheckAndSetDefaults()) + + uploader := eventstest.NewMemoryUploader() + localLog, err := events.NewAuditLog(events.AuditLogConfig{ + DataDir: authServerConfig.Dir, + ServerID: authServerConfig.ClusterName, + Clock: authServerConfig.Clock, + UploadHandler: uploader, + }) + require.NoError(t, err) + authServerConfig.AuditLog = localLog + + as, err := authtest.NewAuthServer(authServerConfig) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, as.Close()) }) + + sessionID := session.NewID() + clusterName := authServerConfig.ClusterName + + // Upload a session that has a start event but no end event. + streamer, err := events.NewProtoStreamer(events.ProtoStreamerConfig{Uploader: uploader}) + require.NoError(t, err) + stream, err := streamer.CreateAuditStream(ctx, sessionID) + require.NoError(t, err) + require.NoError(t, stream.RecordEvent(ctx, eventstest.PrepareEvent(&apievents.SessionStart{ + Metadata: apievents.Metadata{ + Type: events.SessionStartEvent, + Code: events.SessionStartCode, + ClusterName: clusterName, + }, + SessionMetadata: apievents.SessionMetadata{SessionID: sessionID.String()}, + UserMetadata: apievents.UserMetadata{User: "alice", Login: "root"}, + TerminalSize: "80:25", + }))) + require.NoError(t, stream.Complete(ctx)) + + // OnUploadComplete must recover the session end from the stream and emit it. + got, err := as.AuthServer.OnUploadComplete(ctx, sessionID) + require.NoError(t, err) + require.NotNil(t, got) + + sessionEnd, ok := got.(*apievents.SessionEnd) + require.True(t, ok, "expected *apievents.SessionEnd, got %T", got) + require.Equal(t, sessionID.String(), sessionEnd.GetSessionID()) + require.Equal(t, events.SessionEndCode, sessionEnd.Code) + require.True(t, sessionEnd.Interactive) + + // The event must have been emitted to the audit log. + require.EventuallyWithT(t, func(t *assert.CollectT) { + emitted, _, err := localLog.SearchEvents(ctx, events.SearchEventsRequest{ + From: authServerConfig.Clock.Now().Add(-time.Hour), + To: authServerConfig.Clock.Now().Add(time.Hour), + EventTypes: []string{events.SessionEndEvent}, + Limit: 1, + Order: types.EventOrderDescending, + }) + require.NoError(t, err) + require.Len(t, emitted, 1, "recovered session end event must appear in audit log") + require.Equal(t, sessionID.String(), emitted[0].(*apievents.SessionEnd).GetSessionID()) + }, 5*time.Second, 100*time.Millisecond) +} + // TestAPILockedOut tests Auth API when there are locks involved. func TestAPILockedOut(t *testing.T) { t.Parallel() @@ -8178,27 +8250,6 @@ func TestLocalServiceRolesHavePermissionsForUploaderService(t *testing.T) { require.NoError(t, err) }) - t.Run("StreamSessionEvents", func(t *testing.T) { - // use a discard log because we don't care if - // the streaming actually succeeds, we just want to make sure RBAC checks - // pass and allow us to enter the audit log code - s := auth.NewServerWithRoles( - srv.AuthServer, - events.NewDiscardAuditLog(), - *authContext, - ) - - eventC, errC := s.StreamSessionEvents(ctx, "foo", 0) - select { - case err := <-errC: - require.NoError(t, err) - default: - // drain eventC to prevent goroutine leak - for range eventC { - } - } - }) - t.Run("CreateAuditStream", func(t *testing.T) { s := auth.NewServerWithRoles( srv.AuthServer, diff --git a/lib/auth/grpcserver.go b/lib/auth/grpcserver.go index 3de0e77c01e..451c8230590 100644 --- a/lib/auth/grpcserver.go +++ b/lib/auth/grpcserver.go @@ -6516,7 +6516,7 @@ func NewGRPCServer(cfg GRPCServerConfig) (*GRPCServer, error) { Logger: cfg.AuthServer.logger.With(teleport.ComponentKey, teleport.ComponentRecordingEncryption), SessionSummarizerProvider: cfg.APIConfig.AuthServer.sessionSummarizerProvider, RecordingMetadataProvider: cfg.AuthServer.recordingMetadataProvider, - SessionStreamer: cfg.AuthServer, + OnUploadComplete: cfg.AuthServer.OnUploadComplete, }) if err != nil { return nil, trace.Wrap(err) diff --git a/lib/auth/recordingencryption/recordingencryptionv1/service.go b/lib/auth/recordingencryption/recordingencryptionv1/service.go index 6a24d3d0a71..ceabd39cd30 100644 --- a/lib/auth/recordingencryption/recordingencryptionv1/service.go +++ b/lib/auth/recordingencryption/recordingencryptionv1/service.go @@ -27,6 +27,7 @@ import ( "github.com/gravitational/teleport" recordingencryptionv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/recordingencryption/v1" + apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/auth/recordingmetadata" "github.com/gravitational/teleport/lib/auth/summarizer" "github.com/gravitational/teleport/lib/authz" @@ -55,8 +56,9 @@ type ServiceConfig struct { SessionSummarizerProvider *summarizer.SessionSummarizerProvider // RecordingMetadataProvider is a provider of the recording metadata service. RecordingMetadataProvider *recordingmetadata.Provider - // SessionStreamer is a streamer for session events. - SessionStreamer events.SessionStreamer + // OnUploadComplete is called after an upload completes to find or recover the + // session end event. + OnUploadComplete func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) } // NewService returns a new [Service] based on the given [ServiceConfig]. @@ -68,12 +70,12 @@ func NewService(cfg ServiceConfig) (*Service, error) { return nil, trace.BadParameter("uploader is required") case cfg.KeyRotater == nil: return nil, trace.BadParameter("key rotater is required") - case cfg.SessionStreamer == nil: - return nil, trace.BadParameter("session streamer is required") case cfg.RecordingMetadataProvider == nil: return nil, trace.BadParameter("recording metadata provider is required") case cfg.SessionSummarizerProvider == nil: return nil, trace.BadParameter("session summarizer provider is required") + case cfg.OnUploadComplete == nil: + return nil, trace.BadParameter("on upload complete callback is required") } if cfg.Logger == nil { @@ -87,7 +89,7 @@ func NewService(cfg ServiceConfig) (*Service, error) { rotater: cfg.KeyRotater, sessionSummarizerProvider: cfg.SessionSummarizerProvider, recordingMetadataProvider: cfg.RecordingMetadataProvider, - streamer: cfg.SessionStreamer, + onUploadComplete: cfg.OnUploadComplete, }, nil } @@ -99,13 +101,15 @@ type Service struct { logger *slog.Logger uploader events.MultipartUploader rotater KeyRotater - // SessionSummarizerProvider is a provider of the session summarizer service. + // sessionSummarizerProvider is a provider of the session summarizer service. // It can be nil or provide a nil summarizer if summarization is not needed. // The summarizer itself summarizes session recordings. sessionSummarizerProvider *summarizer.SessionSummarizerProvider - // RecordingMetadataProvider is a provider of the recording metadata service. + // recordingMetadataProvider is a provider of the recording metadata service. recordingMetadataProvider *recordingmetadata.Provider - streamer events.SessionStreamer + // onUploadComplete is called after an upload completes to find or recover the + // session end event for post-processing. + onUploadComplete func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) } func streamUploadAsProto(upload events.StreamUpload) *recordingencryptionv1.Upload { @@ -228,7 +232,11 @@ func (s *Service) CompleteUpload(ctx context.Context, req *recordingencryptionv1 return nil, trace.Wrap(err) } - sessionEnd, err := events.FindSessionEndEvent(ctx, s.streamer, upload.SessionID) + if s.onUploadComplete == nil { + return &recordingencryptionv1.CompleteUploadResponse{}, nil + } + + sessionEnd, err := s.onUploadComplete(ctx, upload.SessionID) if err != nil || sessionEnd == nil { return &recordingencryptionv1.CompleteUploadResponse{}, nil } diff --git a/lib/auth/recordingencryption/recordingencryptionv1/service_test.go b/lib/auth/recordingencryption/recordingencryptionv1/service_test.go index 733d8d0585f..8d61de869bc 100644 --- a/lib/auth/recordingencryption/recordingencryptionv1/service_test.go +++ b/lib/auth/recordingencryption/recordingencryptionv1/service_test.go @@ -25,6 +25,7 @@ import ( "github.com/google/uuid" "github.com/gravitational/trace" + "github.com/jonboulle/clockwork" "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" "google.golang.org/protobuf/types/known/timestamppb" @@ -73,9 +74,9 @@ func TestRotateKey(t *testing.T) { Logger: logtest.NewLogger(), Uploader: fakeUploader{}, KeyRotater: rotater, - SessionStreamer: &fakeSessionStreamer{}, RecordingMetadataProvider: recordingmetadata.NewProvider(), SessionSummarizerProvider: summarizer.NewSessionSummarizerProvider(), + OnUploadComplete: func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) { return nil, nil }, } service, err := recordingencryptionv1.NewService(cfg) @@ -120,9 +121,9 @@ func TestCompleteRotation(t *testing.T) { Logger: logtest.NewLogger(), Uploader: fakeUploader{}, KeyRotater: rotater, - SessionStreamer: &fakeSessionStreamer{}, RecordingMetadataProvider: recordingmetadata.NewProvider(), SessionSummarizerProvider: summarizer.NewSessionSummarizerProvider(), + OnUploadComplete: func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) { return nil, nil }, } service, err := recordingencryptionv1.NewService(cfg) @@ -170,9 +171,9 @@ func TestRollbackRotation(t *testing.T) { Logger: logtest.NewLogger(), Uploader: fakeUploader{}, KeyRotater: rotater, - SessionStreamer: &fakeSessionStreamer{}, RecordingMetadataProvider: recordingmetadata.NewProvider(), SessionSummarizerProvider: summarizer.NewSessionSummarizerProvider(), + OnUploadComplete: func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) { return nil, nil }, } service, err := recordingencryptionv1.NewService(cfg) @@ -219,9 +220,9 @@ func TestGetRotationState(t *testing.T) { Logger: logtest.NewLogger(), Uploader: fakeUploader{}, KeyRotater: rotater, - SessionStreamer: &fakeSessionStreamer{}, RecordingMetadataProvider: recordingmetadata.NewProvider(), SessionSummarizerProvider: summarizer.NewSessionSummarizerProvider(), + OnUploadComplete: func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) { return nil, nil }, } service, err := recordingencryptionv1.NewService(cfg) @@ -325,22 +326,6 @@ func (f *fakeKeyRotater) GetRotationState(ctx context.Context) ([]*recordingencr return f.keys, nil } -type fakeSessionStreamer struct{} - -func (f *fakeSessionStreamer) StreamSessionEvents(ctx context.Context, sessionID session.ID, startIndex int64) (chan apievents.AuditEvent, chan error) { - returnChan := make(chan apievents.AuditEvent, 1) - errChan := make(chan error, 1) - close(errChan) - events := eventstest.GenerateTestSession(eventstest.SessionParams{ - UserName: "alice", - SessionID: string(sessionID), - ServerID: "testcluster", - PrintData: []string{"net", "stat"}, - }) - returnChan <- events[len(events)-1] - return returnChan, nil -} - func TestSessionCompleter(t *testing.T) { sessionID := session.ID(uuid.NewString()) @@ -361,9 +346,16 @@ func TestSessionCompleter(t *testing.T) { Logger: logtest.NewLogger(), Uploader: fakeUploader{}, KeyRotater: newFakeKeyRotater(), - SessionStreamer: &fakeSessionStreamer{}, RecordingMetadataProvider: metadataProvider, SessionSummarizerProvider: summarizerProvider, + OnUploadComplete: func(_ context.Context, sid session.ID) (apievents.AuditEvent, error) { + now := time.Now() + return &apievents.SessionEnd{ + SessionMetadata: apievents.SessionMetadata{SessionID: string(sid)}, + StartTime: now.Add(-time.Minute), + EndTime: now, + }, nil + }, } service, err := recordingencryptionv1.NewService(cfg) @@ -411,6 +403,76 @@ func (f *fakeSummarizer) SummarizeWithoutEndEvent(ctx context.Context, sessionID return args.Error(0) } +// TestCompleteUploadRecoversMissingSessionEnd verifies that when an encrypted +// session recording has no session end event (e.g. due to a connection drop), +// CompleteUpload uses the OnUploadComplete callback to recover and emit one. +func TestCompleteUploadRecoversMissingSessionEnd(t *testing.T) { + sessionID := session.ID(uuid.NewString()) + const clusterName = "test-cluster" + + userMeta := apievents.UserMetadata{User: "alice", Login: "root"} + sessionMeta := apievents.SessionMetadata{SessionID: string(sessionID)} + + // Session events without a session end — simulates a dropped connection. + sessionEvents := []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + TerminalSize: "80:25", + }, + &apievents.SessionPrint{ + Metadata: apievents.Metadata{Type: events.SessionPrintEvent}, + }, + } + + auditLog := &eventstest.MockRecorderEmitter{} + streamer := eventstest.NewFakeStreamer(sessionEvents, 0) + + slog := logtest.NewLogger() + + cfg := recordingencryptionv1.ServiceConfig{ + Authorizer: &fakeAuthorizer{}, + Logger: slog, + Uploader: fakeUploader{}, + KeyRotater: newFakeKeyRotater(), + RecordingMetadataProvider: recordingmetadata.NewProvider(), + SessionSummarizerProvider: summarizer.NewSessionSummarizerProvider(), + OnUploadComplete: func(ctx context.Context, sid session.ID) (apievents.AuditEvent, error) { + return events.FindOrRecoverSessionEnd(ctx, events.FindOrRecoverSessionEndConfig{ + ClusterName: clusterName, + Streamer: streamer, + SessionID: sid, + AuditLog: auditLog, + Log: slog, + Clock: clockwork.NewRealClock(), + }) + }, + } + + service, err := recordingencryptionv1.NewService(cfg) + require.NoError(t, err) + + ctx := withAuthCtx(t.Context(), newServiceAuthCtx()) + _, err = service.CompleteUpload(ctx, &recordingencryptionv1pb.CompleteUploadRequest{ + Upload: &recordingencryptionv1pb.Upload{ + SessionId: string(sessionID), + InitiatedAt: timestamppb.Now(), + UploadId: uuid.NewString(), + }, + }) + require.NoError(t, err) + + // The recovered session end event must have been emitted to the audit log. + emitted := auditLog.Events() + require.Len(t, emitted, 1) + sessionEnd, ok := emitted[0].(*apievents.SessionEnd) + require.True(t, ok, "expected *apievents.SessionEnd, got %T", emitted[0]) + require.Equal(t, string(sessionID), sessionEnd.GetSessionID()) + require.Equal(t, userMeta, sessionEnd.UserMetadata) + require.True(t, sessionEnd.Interactive) +} + func newServiceAuthCtx() authz.Context { return authz.Context{ Identity: authz.BuiltinRole{ diff --git a/lib/authz/permissions.go b/lib/authz/permissions.go index 6ae4545e580..3752e06f1af 100644 --- a/lib/authz/permissions.go +++ b/lib/authz/permissions.go @@ -914,8 +914,7 @@ func roleSpecForProxy(clusterName string) types.RoleSpecV6 { types.NewRule(types.KindProxy, services.RW()), types.NewRule(types.KindOIDCRequest, services.RW()), types.NewRule(types.KindSSHSession, services.RW()), - types.NewRule(types.KindSession, services.RO()), - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindSAMLRequest, services.RW()), types.NewRule(types.KindOIDC, services.ReadNoSecrets()), types.NewRule(types.KindSAML, services.ReadNoSecrets()), @@ -1085,8 +1084,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses Rules: []types.Rule{ types.NewRule(types.KindNode, services.RW()), types.NewRule(types.KindSSHSession, services.RW()), - types.NewRule(types.KindSession, services.RO()), - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindProxy, services.RO()), types.NewRule(types.KindCertAuthority, services.ReadNoSecrets()), types.NewRule(types.KindUser, services.RO()), @@ -1122,7 +1120,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses Namespaces: []string{types.Wildcard}, AppLabels: types.Labels{types.Wildcard: []string{types.Wildcard}}, Rules: []types.Rule{ - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindProxy, services.RO()), types.NewRule(types.KindCertAuthority, services.ReadNoSecrets()), types.NewRule(types.KindUser, services.RO()), @@ -1151,7 +1149,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses Namespaces: []string{types.Wildcard}, DatabaseLabels: types.Labels{types.Wildcard: []string{types.Wildcard}}, Rules: []types.Rule{ - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindProxy, services.RO()), types.NewRule(types.KindCertAuthority, services.ReadNoSecrets()), types.NewRule(types.KindUser, services.RO()), @@ -1213,7 +1211,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses types.NewRule(types.KindClusterAuthPreference, services.RO()), types.NewRule(types.KindClusterNetworkingConfig, services.RO()), types.NewRule(types.KindDatabaseServer, services.RO()), - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindKubeServer, services.RO()), types.NewRule(types.KindLock, services.RO()), types.NewRule(types.KindNode, services.RO()), @@ -1284,7 +1282,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses Rules: []types.Rule{ types.NewRule(types.KindKubeServer, services.RW()), types.NewRule(types.KindKubeWaitingContainer, services.RW()), - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindCertAuthority, services.ReadNoSecrets()), types.NewRule(types.KindClusterName, services.RO()), types.NewRule(types.KindClusterAuditConfig, services.RO()), @@ -1309,7 +1307,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses Namespaces: []string{types.Wildcard}, WindowsDesktopLabels: types.Labels{types.Wildcard: []string{types.Wildcard}}, Rules: []types.Rule{ - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindCertAuthority, services.ReadNoSecrets()), types.NewRule(types.KindClusterName, services.RO()), types.NewRule(types.KindClusterAuditConfig, services.RO()), @@ -1333,7 +1331,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses Allow: types.RoleConditions{ Namespaces: []string{types.Wildcard}, Rules: []types.Rule{ - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindCertAuthority, services.ReadNoSecrets()), types.NewRule(types.KindClusterName, services.RO()), types.NewRule(types.KindNamespace, services.RO()), @@ -1366,7 +1364,7 @@ func unscopedDefinitionForBuiltinRole(clusterName string, recConfig readonly.Ses types.NewRule(types.KindClusterName, services.RO()), types.NewRule(types.KindCertAuthority, services.ReadNoSecrets()), types.NewRule(types.KindSemaphore, services.RW()), - types.NewRule(types.KindEvent, services.RW()), + types.NewRule(types.KindEvent, services.WO()), types.NewRule(types.KindAppServer, services.RW()), types.NewRule(types.KindClusterNetworkingConfig, services.RO()), types.NewRule(types.KindUser, services.RW()), diff --git a/lib/events/api.go b/lib/events/api.go index b25d9ecffb9..f96ae927134 100644 --- a/lib/events/api.go +++ b/lib/events/api.go @@ -1135,6 +1135,18 @@ type Streamer interface { ResumeAuditStream(ctx context.Context, sid session.ID, uploadID string) (apievents.Stream, error) } +// StreamerWithCallback extends [Streamer] to allow setting a callback that is +// invoked when a session recording upload completes without a session end event. +type StreamerWithCallback interface { + Streamer + // SetOnUploadComplete registers a callback invoked after a session + // recording upload completes without a session end event. The callback + // may return a session end event or nil if unavailable. + // + // MUST be called before any streams are created. + SetOnUploadComplete(func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error)) +} + // StreamPart represents uploaded stream part type StreamPart struct { // Number is a part number diff --git a/lib/events/auditlog.go b/lib/events/auditlog.go index 2ff72e92d02..6f6a2064e78 100644 --- a/lib/events/auditlog.go +++ b/lib/events/auditlog.go @@ -263,6 +263,10 @@ type AuditLogConfig struct { SessionSummarizerProvider *summarizer.SessionSummarizerProvider // RecordingMetadataProvider provides recording metadata service RecordingMetadataProvider *recordingmetadata.Provider + // OnUploadComplete is called after an encrypted upload completes to find or + // recover the session end event for post-processing. If nil, no + // post-processing is performed. + OnUploadComplete func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) } // CheckAndSetDefaults checks and sets defaults @@ -680,7 +684,10 @@ func (l *AuditLog) UploadEncryptedRecording(ctx context.Context, sessionID strin return trace.Wrap(err, "completing upload") } - sessionEnd, err := FindSessionEndEvent(ctx, l, session.ID(sessionID)) + if l.OnUploadComplete == nil { + return nil + } + sessionEnd, err := l.OnUploadComplete(ctx, upload.SessionID) if err != nil || sessionEnd == nil { return nil } @@ -698,6 +705,12 @@ func (l *AuditLog) UploadEncryptedRecording(ctx context.Context, sessionID strin return nil } +// SetOnUploadComplete sets the callback to be invoked after an encrypted upload +// completes. It must be called before any uploads are processed. +func (l *AuditLog) SetOnUploadComplete(fn func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error)) { + l.OnUploadComplete = fn +} + // getLocalLog returns the local (file based) AuditLogger. func (l *AuditLog) getLocalLog() AuditLogger { l.RLock() diff --git a/lib/events/auditlog_test.go b/lib/events/auditlog_test.go index a47c4551b11..642aef20ee4 100644 --- a/lib/events/auditlog_test.go +++ b/lib/events/auditlog_test.go @@ -410,6 +410,14 @@ func TestCallingSummarizerMetadata(t *testing.T) { UploadHandler: uploader, SessionSummarizerProvider: summarizerProvider, RecordingMetadataProvider: metadataProvider, + OnUploadComplete: func(_ context.Context, sid session.ID) (apievents.AuditEvent, error) { + now := time.Now() + return &apievents.SessionEnd{ + SessionMetadata: apievents.SessionMetadata{SessionID: string(sid)}, + StartTime: now.Add(-time.Minute), + EndTime: now, + }, nil + }, }) require.NoError(t, err) defer alog.Close() diff --git a/lib/events/complete.go b/lib/events/complete.go index 30353f9c3ad..82ce31b2cd6 100644 --- a/lib/events/complete.go +++ b/lib/events/complete.go @@ -21,7 +21,6 @@ package events import ( "cmp" "context" - "fmt" "log/slog" "time" @@ -32,14 +31,12 @@ import ( "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/api/types/events" - apiutils "github.com/gravitational/teleport/api/utils" + apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/auth/recordingmetadata" "github.com/gravitational/teleport/lib/auth/summarizer" "github.com/gravitational/teleport/lib/observability/metrics" "github.com/gravitational/teleport/lib/services" - "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/teleport/lib/utils/interval" ) @@ -76,6 +73,9 @@ type UploadCompleterConfig struct { SessionSummarizerProvider *summarizer.SessionSummarizerProvider // RecordingMetadataProvider is a provider of the recording metadata service. RecordingMetadataProvider *recordingmetadata.Provider + // EnsureSessionEndEvent determines whether or not the UploadCompleter should + // detect missing session end events and attempt to emit them. + EnsureSessionEndEvent bool } // CheckAndSetDefaults checks and sets default values @@ -315,19 +315,21 @@ func (u *UploadCompleter) CheckUploads(ctx context.Context) error { // This is necessary because we'll need to download the session in order to // enumerate its events, and the S3 API takes a little while after the upload // is completed before version metadata becomes available. - go func() { - select { - case <-ctx.Done(): - return - case <-u.cfg.Clock.After(2 * time.Minute): - log.DebugContext(ctx, "checking for session end event") - if err := u.ensureSessionEndEvent(ctx, uploadData); err != nil { - log.WarnContext(ctx, "failed to ensure session end event", "error", err) + if u.cfg.EnsureSessionEndEvent { + go func() { + select { + case <-ctx.Done(): + return + case <-u.cfg.Clock.After(2 * time.Minute): + log.DebugContext(ctx, "checking for session end event") + if err := u.ensureSessionEndEvent(ctx, uploadData); err != nil { + log.WarnContext(ctx, "failed to ensure session end event", "error", err) + } } - } - }() - session := &events.SessionUpload{ - Metadata: events.Metadata{ + }() + } + session := &apievents.SessionUpload{ + Metadata: apievents.Metadata{ Type: SessionUploadEvent, Code: SessionUploadCode, Time: u.cfg.Clock.Now().UTC(), @@ -335,7 +337,7 @@ func (u *UploadCompleter) CheckUploads(ctx context.Context) error { Index: SessionUploadIndex, ClusterName: u.cfg.ClusterName, }, - SessionMetadata: events.SessionMetadata{ + SessionMetadata: apievents.SessionMetadata{ SessionID: string(uploadData.SessionID), }, SessionURL: uploadData.URL, @@ -349,145 +351,46 @@ func (u *UploadCompleter) CheckUploads(ctx context.Context) error { } func (u *UploadCompleter) ensureSessionEndEvent(ctx context.Context, uploadData UploadMetadata) error { - // at this point, we don't know whether we'll need to emit a session.end or a - // windows.desktop.session.end, but as soon as we see the session start we'll - // be able to start filling in the details - var sshSessionEnd events.SessionEnd - var desktopSessionEnd events.WindowsDesktopSessionEnd - - // We use the streaming events API to search through the session events, because it works - // for both Desktop and SSH sessions - var lastEvent events.AuditEvent - var startTime time.Time - var sessionType recordingmetadata.SessionType - var isPTYSession bool - ctx, cancel := context.WithCancel(ctx) - defer cancel() - evts, errors := u.cfg.AuditLog.StreamSessionEvents(ctx, uploadData.SessionID, 0) - -loop: - for { - select { - case evt, more := <-evts: - if !more { - break loop - } - - lastEvent = evt - - switch e := evt.(type) { - // Return if session end event already exists - case *events.SessionEnd, *events.WindowsDesktopSessionEnd: - return nil - - case *events.WindowsDesktopSessionStart: - startTime = e.Time - desktopSessionEnd.Type = WindowsDesktopSessionEndEvent - desktopSessionEnd.Code = DesktopSessionEndCode - desktopSessionEnd.ClusterName = e.ClusterName - desktopSessionEnd.StartTime = e.Time - desktopSessionEnd.Participants = append(desktopSessionEnd.Participants, transformedUsername(e.UserMetadata, u.cfg.ClusterName)) - desktopSessionEnd.Recorded = true - desktopSessionEnd.UserMetadata = e.UserMetadata - desktopSessionEnd.SessionMetadata = e.SessionMetadata - desktopSessionEnd.WindowsDesktopService = e.WindowsDesktopService - desktopSessionEnd.Domain = e.Domain - desktopSessionEnd.DesktopAddr = e.DesktopAddr - desktopSessionEnd.DesktopLabels = e.DesktopLabels - desktopSessionEnd.DesktopName = fmt.Sprintf("%v (recovered)", e.DesktopName) - - case *events.SessionStart: - sessionType = recordingmetadata.SessionTypeTTY - isPTYSession = true - startTime = e.Time - sshSessionEnd.Type = SessionEndEvent - sshSessionEnd.Code = SessionEndCode - sshSessionEnd.ClusterName = e.ClusterName - sshSessionEnd.StartTime = e.Time - sshSessionEnd.UserMetadata = e.UserMetadata - sshSessionEnd.SessionMetadata = e.SessionMetadata - sshSessionEnd.ServerMetadata = e.ServerMetadata - sshSessionEnd.ConnectionMetadata = e.ConnectionMetadata - sshSessionEnd.KubernetesClusterMetadata = e.KubernetesClusterMetadata - sshSessionEnd.KubernetesPodMetadata = e.KubernetesPodMetadata - sshSessionEnd.InitialCommand = e.InitialCommand - sshSessionEnd.SessionRecording = e.SessionRecording - sshSessionEnd.Interactive = e.TerminalSize != "" - sshSessionEnd.Participants = append(sshSessionEnd.Participants, transformedUsername(e.UserMetadata, u.cfg.ClusterName)) - - case *events.DatabaseSessionStart: - startTime = e.Time - - case *events.SessionJoin: - sshSessionEnd.Participants = append(sshSessionEnd.Participants, transformedUsername(e.UserMetadata, u.cfg.ClusterName)) - } - - case err := <-errors: - return trace.Wrap(err) - case <-ctx.Done(): - return ctx.Err() - } - } - - if lastEvent == nil { - return trace.Errorf("could not find any events for session %v", uploadData.SessionID) - } - - sshSessionEnd.Participants = apiutils.Deduplicate(sshSessionEnd.Participants) - sshSessionEnd.EndTime = lastEvent.GetTime() - desktopSessionEnd.EndTime = lastEvent.GetTime() - - var sessionEndEvent events.AuditEvent - switch { - case sshSessionEnd.Code != "": - sessionEndEvent = &sshSessionEnd - case desktopSessionEnd.Code != "": - sessionEndEvent = &desktopSessionEnd - default: - return trace.BadParameter("invalid session, could not find session start") - } - - u.log.InfoContext(ctx, "emitting event for completed session", - "event_type", sessionEndEvent.GetType(), - "event_code", sessionEndEvent.GetCode(), - "session_id", uploadData.SessionID, + sessionEndEvent, err := FindOrRecoverSessionEnd( + ctx, + FindOrRecoverSessionEndConfig{ + ClusterName: u.cfg.ClusterName, + Streamer: u.cfg.AuditLog, + SessionID: uploadData.SessionID, + AuditLog: u.cfg.AuditLog, + Log: u.log, + Clock: u.cfg.Clock, + }, ) - - sessionEndEvent.SetTime(lastEvent.GetTime()) - - // Check and set event fields - if err := checkAndSetEventFields(sessionEndEvent, u.cfg.Clock, utils.NewRealUID(), sessionEndEvent.GetClusterName()); err != nil { + if err != nil { return trace.Wrap(err) } - if err := u.cfg.AuditLog.EmitAuditEvent(ctx, sessionEndEvent); err != nil { - return trace.Wrap(err) - } - - if !isPTYSession { - return nil - } // For PTY sessions, process recording metadata and summarization. recordingMetadata := u.cfg.RecordingMetadataProvider.Service() - if !startTime.IsZero() && !sessionEndEvent.GetTime().IsZero() { - duration := sessionEndEvent.GetTime().Sub(startTime) + if duration, sessionType := isPTYSession(sessionEndEvent); duration > 0 { if err := recordingMetadata.ProcessSessionRecording(ctx, uploadData.SessionID, sessionType, duration); err != nil { slog.WarnContext(ctx, "Failed to process session recording metadata", "error", err) } - } else { + } else if sessionType == recordingmetadata.SessionTypeTTY { slog.WarnContext(ctx, "Session start or end time is not set, skipping recording metadata processing") } summarizer := u.cfg.SessionSummarizerProvider.SessionSummarizer() - if err := summarizer.SummarizeSSH(ctx, &sshSessionEnd); err != nil { + switch o := sessionEndEvent.(type) { + case *apievents.SessionEnd: + err = summarizer.SummarizeSSH(ctx, o) + case *apievents.DatabaseSessionEnd: + err = summarizer.SummarizeDatabase(ctx, o) + } + if err != nil { slog.WarnContext(ctx, "Failed to summarize upload", "error", err) return trace.Wrap(err) } - return nil } -func transformedUsername(u events.UserMetadata, localCluster string) string { +func transformedUsername(u apievents.UserMetadata, localCluster string) string { return services.UsernameForCluster( services.UsernameForClusterConfig{ User: u.User, @@ -496,3 +399,28 @@ func transformedUsername(u events.UserMetadata, localCluster string) string { }, ) } + +// isPTYSession returns the session duration and true if the event is an +// interactive (PTY) SSH session end. Returns (0, false) for non-SSH sessions. +// Returns (-1, true) if the event is a PTY session but the start or end time +// is missing. +func isPTYSession(sessionEnd apievents.AuditEvent) (time.Duration, recordingmetadata.SessionType) { + switch evt := sessionEnd.(type) { + case *apievents.SessionEnd: + if evt.EndTime.IsZero() || evt.StartTime.IsZero() { + return -1, recordingmetadata.SessionTypeTTY + } + return evt.EndTime.Sub(evt.StartTime), recordingmetadata.SessionTypeTTY + case *apievents.DatabaseSessionEnd: + if evt.EndTime.IsZero() || evt.StartTime.IsZero() { + return -1, recordingmetadata.SessionTypeUnspecified + } + return evt.EndTime.Sub(evt.StartTime), recordingmetadata.SessionTypeUnspecified + case *apievents.WindowsDesktopSessionEnd: + if evt.EndTime.IsZero() || evt.StartTime.IsZero() { + return -1, recordingmetadata.SessionTypeUnspecified + } + return evt.EndTime.Sub(evt.StartTime), recordingmetadata.SessionTypeUnspecified + } + return 0, recordingmetadata.SessionTypeUnspecified +} diff --git a/lib/events/complete_test.go b/lib/events/complete_test.go index ff76a34f117..633e1668c03 100644 --- a/lib/events/complete_test.go +++ b/lib/events/complete_test.go @@ -24,6 +24,7 @@ import ( "fmt" "strings" "testing" + "testing/synctest" "time" "github.com/gravitational/trace" @@ -177,63 +178,67 @@ func TestUploadCompleterAcquiresSemaphore(t *testing.T) { // that are completed. func TestUploadCompleterEmitsSessionEnd(t *testing.T) { for _, test := range []struct { - startEvent apievents.AuditEvent - endEventType string + startEvent apievents.AuditEvent + endEventType string + ensureSessionEndEvent bool }{ - {&apievents.SessionStart{}, events.SessionEndEvent}, - {&apievents.WindowsDesktopSessionStart{}, events.WindowsDesktopSessionEndEvent}, + {&apievents.SessionStart{}, events.SessionEndEvent, true}, + {&apievents.WindowsDesktopSessionStart{}, events.WindowsDesktopSessionEndEvent, true}, + {&apievents.SessionStart{}, events.SessionEndEvent, false}, } { - t.Run(test.endEventType, func(t *testing.T) { - clock := clockwork.NewFakeClock() - mu := eventstest.NewMemoryUploader() - mu.Clock = clock - startTime := clock.Now().UTC() - endTime := startTime.Add(2 * time.Minute) + t.Run(fmt.Sprintf("%s ensure end event %t", test.endEventType, test.ensureSessionEndEvent), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + clock := clockwork.NewFakeClock() + mu := eventstest.NewMemoryUploader() + startTime := clock.Now().UTC() + endTime := startTime.Add(2 * time.Minute) - test.startEvent.SetTime(startTime) + test.startEvent.SetTime(startTime) - log := &eventstest.MockAuditLog{ - Emitter: &eventstest.MockRecorderEmitter{}, - SessionEvents: []apievents.AuditEvent{ - test.startEvent, - &apievents.SessionPrint{Metadata: apievents.Metadata{Time: endTime}}, - }, - } + log := &eventstest.MockAuditLog{ + Emitter: &eventstest.MockRecorderEmitter{}, + SessionEvents: []apievents.AuditEvent{ + test.startEvent, + &apievents.SessionPrint{Metadata: apievents.Metadata{Time: endTime}}, + }, + } - uc, err := events.NewUploadCompleter(events.UploadCompleterConfig{ - Uploader: mu, - AuditLog: log, - Clock: clock, - SessionTracker: &mockSessionTrackerService{}, - ClusterName: "teleport-cluster", - GracePeriod: -1, + uc, err := events.NewUploadCompleter(events.UploadCompleterConfig{ + Uploader: mu, + AuditLog: log, + SessionTracker: &mockSessionTrackerService{}, + ClusterName: "teleport-cluster", + GracePeriod: -1, + EnsureSessionEndEvent: test.ensureSessionEndEvent, + }) + require.NoError(t, err) + + upload, err := mu.CreateUpload(context.Background(), session.NewID()) + require.NoError(t, err) + + // session end events are only emitted if there's at least one + // part to be uploaded, so create that here + _, err = mu.UploadPart(context.Background(), *upload, 0, strings.NewReader("part")) + require.NoError(t, err) + + err = uc.CheckUploads(context.Background()) + require.NoError(t, err) + + time.Sleep(3 * time.Minute) + synctest.Wait() + synctest.Wait() + + if len(log.Emitter.Events()) == 2 && !test.ensureSessionEndEvent { + require.FailNow(t, "should only have emitted 1 session upload event") + } + + require.IsType(t, &apievents.SessionUpload{}, log.Emitter.Events()[0]) + require.Equal(t, startTime, log.Emitter.Events()[0].GetTime()) + if test.ensureSessionEndEvent { + require.Equal(t, test.endEventType, log.Emitter.Events()[1].GetType()) + require.Equal(t, endTime, log.Emitter.Events()[1].GetTime()) + } }) - require.NoError(t, err) - - upload, err := mu.CreateUpload(context.Background(), session.NewID()) - require.NoError(t, err) - - // session end events are only emitted if there's at least one - // part to be uploaded, so create that here - _, err = mu.UploadPart(context.Background(), *upload, 0, strings.NewReader("part")) - require.NoError(t, err) - - err = uc.CheckUploads(context.Background()) - require.NoError(t, err) - - // advance the clock to force the asynchronous session end event emission - clock.BlockUntil(1) - clock.Advance(3 * time.Minute) - - // expect two events - a session end and a session upload - // the session end is done asynchronously, so wait for that - require.Eventually(t, func() bool { return len(log.Emitter.Events()) == 2 }, 5*time.Second, 1*time.Second, - "should have emitted 2 events, but only got %d", len(log.Emitter.Events())) - - require.IsType(t, &apievents.SessionUpload{}, log.Emitter.Events()[0]) - require.Equal(t, startTime, log.Emitter.Events()[0].GetTime()) - require.Equal(t, test.endEventType, log.Emitter.Events()[1].GetType()) - require.Equal(t, endTime, log.Emitter.Events()[1].GetTime()) }) } } diff --git a/lib/events/discard.go b/lib/events/discard.go index b9dab764581..c87f6952870 100644 --- a/lib/events/discard.go +++ b/lib/events/discard.go @@ -186,6 +186,10 @@ func (*DiscardStreamer) ResumeAuditStream(ctx context.Context, sid session.ID, u return NewDiscardRecorder(), nil } +func (*DiscardStreamer) SetOnUploadComplete(func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error)) { + // no-op +} + // NoOpPreparer is a SessionEventPreparer that doesn't change events type NoOpPreparer struct{} diff --git a/lib/events/sessionend.go b/lib/events/sessionend.go index f8533dd6798..70652b1865f 100644 --- a/lib/events/sessionend.go +++ b/lib/events/sessionend.go @@ -20,11 +20,16 @@ package events import ( "context" + "fmt" + "log/slog" "github.com/gravitational/trace" + "github.com/jonboulle/clockwork" apievents "github.com/gravitational/teleport/api/types/events" + apiutils "github.com/gravitational/teleport/api/utils" "github.com/gravitational/teleport/lib/session" + "github.com/gravitational/teleport/lib/utils" ) // FindSessionEndEvent streams session events to find the session end event for the given session ID. @@ -73,3 +78,205 @@ func FindSessionEndEvent(ctx context.Context, streamer SessionStreamer, sessionI } } } + +// FindOrRecoverSessionEndConfig holds the configuration for FindOrRecoverSessionEnd. +type FindOrRecoverSessionEndConfig struct { + // ClusterName is the name of the cluster where the session took place. + ClusterName string + // Streamer is used to stream session events. + Streamer SessionStreamer + // SessionID is the ID of the session to recover. + SessionID session.ID + // AuditLog is the emitter used to emit the recovered session end event. + AuditLog apievents.Emitter + // Log is the logger. + Log *slog.Logger + // Clock is used to timestamp the recovered event. + Clock clockwork.Clock +} + +// FindOrRecoverSessionEnd streams the events for the given session and returns +// the session end event. +// +// If a session end event already exists in the stream, it is returned directly. +// Otherwise, the function reconstructs (recovers) a session end event from the +// session start event and any additional events found in the stream, emits it to +// the audit log, and returns it. +// +// Supported session types are: +// - SSH / Kubernetes (session.start -> session.end) +// - Windows Desktop (windows.desktop.session.start -> windows.desktop.session.end) +// - Database (db.session.start -> db.session.end) +// - Application (app.session.start -> app.session.end) +// - MCP (mcp.session.start -> mcp.session.end) +// +// An error is returned if no events are found for the session, if the session +// start event cannot be identified, or if emitting the recovered event fails. +func FindOrRecoverSessionEnd(ctx context.Context, cfg FindOrRecoverSessionEndConfig) (apievents.AuditEvent, error) { + if err := validateFindOrRecoverSessionEndConfig(cfg); err != nil { + return nil, trace.Wrap(err) + } + // at this point, we don't know which session type we're dealing with, but as + // soon as we see the session start we'll be able to start filling in the details + var sshSessionEnd apievents.SessionEnd + var desktopSessionEnd apievents.WindowsDesktopSessionEnd + var dbSessionEnd apievents.DatabaseSessionEnd + var appSessionEnd apievents.AppSessionEnd + var mcpSessionEnd apievents.MCPSessionEnd + + // We use the streaming events API to search through the session events, because it works + // for all session types + var lastEvent apievents.AuditEvent + ctx, cancel := context.WithCancel(ctx) + defer cancel() + evts, errors := cfg.Streamer.StreamSessionEvents(ctx, cfg.SessionID, 0) + +loop: + for { + select { + case evt, more := <-evts: + if !more { + break loop + } + + lastEvent = evt + + switch e := evt.(type) { + // Return if session end event already exists + case *apievents.SessionEnd, *apievents.WindowsDesktopSessionEnd, + *apievents.DatabaseSessionEnd, *apievents.AppSessionEnd, *apievents.MCPSessionEnd: + return e, nil + + case *apievents.WindowsDesktopSessionStart: + desktopSessionEnd.Type = WindowsDesktopSessionEndEvent + desktopSessionEnd.Code = DesktopSessionEndCode + desktopSessionEnd.ClusterName = e.ClusterName + desktopSessionEnd.StartTime = e.Time + desktopSessionEnd.Participants = append(desktopSessionEnd.Participants, transformedUsername(e.UserMetadata, cfg.ClusterName)) + desktopSessionEnd.Recorded = true + desktopSessionEnd.UserMetadata = e.UserMetadata + desktopSessionEnd.SessionMetadata = e.SessionMetadata + desktopSessionEnd.WindowsDesktopService = e.WindowsDesktopService + desktopSessionEnd.Domain = e.Domain + desktopSessionEnd.DesktopAddr = e.DesktopAddr + desktopSessionEnd.DesktopLabels = e.DesktopLabels + desktopSessionEnd.DesktopName = fmt.Sprintf("%v (recovered)", e.DesktopName) + + case *apievents.SessionStart: + sshSessionEnd.Type = SessionEndEvent + sshSessionEnd.Code = SessionEndCode + sshSessionEnd.ClusterName = e.ClusterName + sshSessionEnd.StartTime = e.Time + sshSessionEnd.UserMetadata = e.UserMetadata + sshSessionEnd.SessionMetadata = e.SessionMetadata + sshSessionEnd.ServerMetadata = e.ServerMetadata + sshSessionEnd.ConnectionMetadata = e.ConnectionMetadata + sshSessionEnd.KubernetesClusterMetadata = e.KubernetesClusterMetadata + sshSessionEnd.KubernetesPodMetadata = e.KubernetesPodMetadata + sshSessionEnd.InitialCommand = e.InitialCommand + sshSessionEnd.SessionRecording = e.SessionRecording + sshSessionEnd.Interactive = e.TerminalSize != "" + sshSessionEnd.Participants = append(sshSessionEnd.Participants, transformedUsername(e.UserMetadata, cfg.ClusterName)) + + case *apievents.SessionJoin: + sshSessionEnd.Participants = append(sshSessionEnd.Participants, transformedUsername(e.UserMetadata, cfg.ClusterName)) + + case *apievents.DatabaseSessionStart: + dbSessionEnd.Type = DatabaseSessionEndEvent + dbSessionEnd.Code = DatabaseSessionEndCode + dbSessionEnd.ClusterName = e.ClusterName + dbSessionEnd.StartTime = e.Time + dbSessionEnd.UserMetadata = e.UserMetadata + dbSessionEnd.SessionMetadata = e.SessionMetadata + dbSessionEnd.DatabaseMetadata = e.DatabaseMetadata + dbSessionEnd.ConnectionMetadata = e.ConnectionMetadata + dbSessionEnd.Participants = append(dbSessionEnd.Participants, transformedUsername(e.UserMetadata, cfg.ClusterName)) + + case *apievents.AppSessionStart: + appSessionEnd.Type = AppSessionEndEvent + appSessionEnd.Code = AppSessionEndCode + appSessionEnd.ClusterName = e.ClusterName + appSessionEnd.UserMetadata = e.UserMetadata + appSessionEnd.SessionMetadata = e.SessionMetadata + appSessionEnd.ServerMetadata = e.ServerMetadata + appSessionEnd.ConnectionMetadata = e.ConnectionMetadata + appSessionEnd.AppMetadata = e.AppMetadata + + case *apievents.MCPSessionStart: + mcpSessionEnd.Type = MCPSessionEndEvent + mcpSessionEnd.Code = MCPSessionEndCode + mcpSessionEnd.ClusterName = e.ClusterName + mcpSessionEnd.UserMetadata = e.UserMetadata + mcpSessionEnd.SessionMetadata = e.SessionMetadata + mcpSessionEnd.ServerMetadata = e.ServerMetadata + mcpSessionEnd.ConnectionMetadata = e.ConnectionMetadata + mcpSessionEnd.AppMetadata = e.AppMetadata + } + + case err := <-errors: + return nil, trace.Wrap(err) + case <-ctx.Done(): + return nil, ctx.Err() + } + } + + if lastEvent == nil { + return nil, trace.Errorf("could not find any events for session %v", cfg.SessionID) + } + + sshSessionEnd.Participants = apiutils.Deduplicate(sshSessionEnd.Participants) + sshSessionEnd.EndTime = lastEvent.GetTime() + desktopSessionEnd.EndTime = lastEvent.GetTime() + dbSessionEnd.EndTime = lastEvent.GetTime() + + var sessionEndEvent apievents.AuditEvent + switch { + case sshSessionEnd.Code != "": + sessionEndEvent = &sshSessionEnd + case desktopSessionEnd.Code != "": + sessionEndEvent = &desktopSessionEnd + case dbSessionEnd.Code != "": + sessionEndEvent = &dbSessionEnd + case appSessionEnd.Code != "": + sessionEndEvent = &appSessionEnd + case mcpSessionEnd.Code != "": + sessionEndEvent = &mcpSessionEnd + default: + return nil, trace.BadParameter("invalid session, could not find session start") + } + + cfg.Log.InfoContext(ctx, "emitting event for completed session", + "event_type", sessionEndEvent.GetType(), + "event_code", sessionEndEvent.GetCode(), + "session_id", cfg.SessionID, + ) + + sessionEndEvent.SetTime(lastEvent.GetTime()) + + // Check and set event fields + if err := checkAndSetEventFields(sessionEndEvent, cfg.Clock, utils.NewRealUID(), sessionEndEvent.GetClusterName()); err != nil { + return nil, trace.Wrap(err) + } + if err := cfg.AuditLog.EmitAuditEvent(ctx, sessionEndEvent); err != nil { + return nil, trace.Wrap(err) + } + return sessionEndEvent, nil +} + +func validateFindOrRecoverSessionEndConfig(cfg FindOrRecoverSessionEndConfig) error { + switch { + case cfg.ClusterName == "": + return trace.BadParameter("ClusterName is required") + case cfg.Streamer == nil: + return trace.BadParameter("Streamer is required") + case cfg.SessionID == "": + return trace.BadParameter("SessionID is required") + case cfg.AuditLog == nil: + return trace.BadParameter("AuditLog is required") + case cfg.Log == nil: + return trace.BadParameter("Log is required") + case cfg.Clock == nil: + return trace.BadParameter("Clock is required") + } + return nil +} diff --git a/lib/events/sessionend_test.go b/lib/events/sessionend_test.go index f83de86f2e0..231b21c0460 100644 --- a/lib/events/sessionend_test.go +++ b/lib/events/sessionend_test.go @@ -19,8 +19,12 @@ package events_test import ( + "log/slog" "testing" + "time" + "github.com/jonboulle/clockwork" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" apievents "github.com/gravitational/teleport/api/types/events" @@ -129,3 +133,375 @@ func TestFindSessionEndEvent(t *testing.T) { }) } } + +func TestFindOrRecoverSessionEnd(t *testing.T) { + const clusterName = "test-cluster" + sessionID := session.NewID() + clock := clockwork.NewFakeClock() + + userMeta := apievents.UserMetadata{User: "alice", Login: "root"} + sessionMeta := apievents.SessionMetadata{SessionID: string(sessionID)} + serverMeta := apievents.ServerMetadata{ServerID: "srv-1", ServerHostname: "host-1"} + connMeta := apievents.ConnectionMetadata{LocalAddr: "127.0.0.1:3022", RemoteAddr: "10.0.0.1:9999"} + dbMeta := apievents.DatabaseMetadata{DatabaseName: "mydb", DatabaseUser: "admin", DatabaseProtocol: "postgres"} + appMeta := apievents.AppMetadata{AppName: "myapp", AppURI: "http://app.local"} + + startTime := clock.Now().UTC() + lastTime := startTime.Add(time.Minute) + + makeConfig := func(streamer events.SessionStreamer, emitter apievents.Emitter) events.FindOrRecoverSessionEndConfig { + return events.FindOrRecoverSessionEndConfig{ + ClusterName: clusterName, + Streamer: streamer, + SessionID: sessionID, + AuditLog: emitter, + Log: slog.Default(), + Clock: clock, + } + } + + t.Run("validation", func(t *testing.T) { + base := makeConfig(eventstest.NewFakeStreamer(nil, 0), &eventstest.MockRecorderEmitter{}) + tests := []struct { + name string + mutate func(*events.FindOrRecoverSessionEndConfig) + }{ + {"missing ClusterName", func(c *events.FindOrRecoverSessionEndConfig) { c.ClusterName = "" }}, + {"missing Streamer", func(c *events.FindOrRecoverSessionEndConfig) { c.Streamer = nil }}, + {"missing SessionID", func(c *events.FindOrRecoverSessionEndConfig) { c.SessionID = "" }}, + {"missing AuditLog", func(c *events.FindOrRecoverSessionEndConfig) { c.AuditLog = nil }}, + {"missing Log", func(c *events.FindOrRecoverSessionEndConfig) { c.Log = nil }}, + {"missing Clock", func(c *events.FindOrRecoverSessionEndConfig) { c.Clock = nil }}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := base + tt.mutate(&c) + _, err := events.FindOrRecoverSessionEnd(t.Context(), c) + require.Error(t, err) + }) + } + }) + + type checkFn func(t *testing.T, gotEnd apievents.AuditEvent, emitted []apievents.AuditEvent) + + // alreadyExists returns a checkFn that asserts the exact start/end pair is + // returned and that no event was emitted (because the end already existed). + alreadyExists := func(wantEnd apievents.AuditEvent) checkFn { + return func(t *testing.T, gotEnd apievents.AuditEvent, emitted []apievents.AuditEvent) { + t.Helper() + assert.Equal(t, wantEnd, gotEnd) + assert.Empty(t, emitted, "expected no event to be emitted when end already exists") + } + } + + tests := []struct { + name string + evts []apievents.AuditEvent + wantErr bool + check checkFn + }{ + { + name: "no events", + evts: nil, + wantErr: true, + }, + { + name: "no session start event", + evts: []apievents.AuditEvent{ + &apievents.SessionPrint{Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: lastTime}}, + }, + wantErr: true, + }, + { + name: "SSH/end already exists", + evts: []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + }, + &apievents.SessionEnd{ + Metadata: apievents.Metadata{Type: events.SessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + }, + }, + check: alreadyExists( + &apievents.SessionEnd{ + Metadata: apievents.Metadata{Type: events.SessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + }, + ), + }, + { + name: "SSH/end recovered", + evts: []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + ServerMetadata: serverMeta, + ConnectionMetadata: connMeta, + TerminalSize: "80:25", + }, + &apievents.SessionPrint{Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: lastTime}}, + }, + check: func(t *testing.T, gotEnd apievents.AuditEvent, emitted []apievents.AuditEvent) { + t.Helper() + recovered, ok := gotEnd.(*apievents.SessionEnd) + require.True(t, ok) + assert.Equal(t, events.SessionEndEvent, recovered.Type) + assert.Equal(t, events.SessionEndCode, recovered.Code) + assert.Equal(t, userMeta, recovered.UserMetadata) + assert.Equal(t, sessionMeta, recovered.SessionMetadata) + assert.Equal(t, serverMeta, recovered.ServerMetadata) + assert.Equal(t, connMeta, recovered.ConnectionMetadata) + assert.True(t, recovered.Interactive) + assert.Equal(t, lastTime, recovered.EndTime) + assert.Len(t, emitted, 1) + }, + }, + { + name: "SSH/participants deduplicated", + evts: []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + }, + &apievents.SessionJoin{ + Metadata: apievents.Metadata{Type: events.SessionJoinEvent, Time: startTime.Add(time.Second)}, + UserMetadata: userMeta, // same user — should be deduplicated + }, + &apievents.SessionPrint{Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: lastTime}}, + }, + check: func(t *testing.T, gotEnd apievents.AuditEvent, _ []apievents.AuditEvent) { + t.Helper() + recovered, ok := gotEnd.(*apievents.SessionEnd) + require.True(t, ok) + assert.Len(t, recovered.Participants, 1) + }, + }, + { + name: "Windows Desktop/end already exists", + evts: []apievents.AuditEvent{ + &apievents.WindowsDesktopSessionStart{ + Metadata: apievents.Metadata{Type: events.WindowsDesktopSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + }, + &apievents.WindowsDesktopSessionEnd{ + Metadata: apievents.Metadata{Type: events.WindowsDesktopSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + }, + }, + check: alreadyExists( + &apievents.WindowsDesktopSessionEnd{ + Metadata: apievents.Metadata{Type: events.WindowsDesktopSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + }, + ), + }, + { + name: "Windows Desktop/end recovered", + evts: []apievents.AuditEvent{ + &apievents.WindowsDesktopSessionStart{ + Metadata: apievents.Metadata{Type: events.WindowsDesktopSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + DesktopName: "mydesktop", + Domain: "CORP", + }, + &apievents.SessionPrint{Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: lastTime}}, + }, + check: func(t *testing.T, gotEnd apievents.AuditEvent, emitted []apievents.AuditEvent) { + t.Helper() + recovered, ok := gotEnd.(*apievents.WindowsDesktopSessionEnd) + require.True(t, ok) + assert.Equal(t, events.WindowsDesktopSessionEndEvent, recovered.Type) + assert.Equal(t, events.DesktopSessionEndCode, recovered.Code) + assert.Equal(t, userMeta, recovered.UserMetadata) + assert.Equal(t, sessionMeta, recovered.SessionMetadata) + assert.Equal(t, "mydesktop (recovered)", recovered.DesktopName) + assert.Equal(t, "CORP", recovered.Domain) + assert.True(t, recovered.Recorded) + assert.Equal(t, lastTime, recovered.EndTime) + assert.Len(t, emitted, 1) + }, + }, + { + name: "Database/end already exists", + evts: []apievents.AuditEvent{ + &apievents.DatabaseSessionStart{ + Metadata: apievents.Metadata{Type: events.DatabaseSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + DatabaseMetadata: dbMeta, + }, + &apievents.DatabaseSessionEnd{ + Metadata: apievents.Metadata{Type: events.DatabaseSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + DatabaseMetadata: dbMeta, + }, + }, + check: alreadyExists( + &apievents.DatabaseSessionEnd{ + Metadata: apievents.Metadata{Type: events.DatabaseSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + DatabaseMetadata: dbMeta, + }, + ), + }, + { + name: "Database/end recovered", + evts: []apievents.AuditEvent{ + &apievents.DatabaseSessionStart{ + Metadata: apievents.Metadata{Type: events.DatabaseSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + DatabaseMetadata: dbMeta, + ConnectionMetadata: connMeta, + }, + &apievents.DatabaseSessionQuery{Metadata: apievents.Metadata{Type: events.DatabaseSessionQueryEvent, Time: lastTime}}, + }, + check: func(t *testing.T, gotEnd apievents.AuditEvent, emitted []apievents.AuditEvent) { + t.Helper() + recovered, ok := gotEnd.(*apievents.DatabaseSessionEnd) + require.True(t, ok) + assert.Equal(t, events.DatabaseSessionEndEvent, recovered.Type) + assert.Equal(t, events.DatabaseSessionEndCode, recovered.Code) + assert.Equal(t, userMeta, recovered.UserMetadata) + assert.Equal(t, sessionMeta, recovered.SessionMetadata) + assert.Equal(t, dbMeta, recovered.DatabaseMetadata) + assert.Equal(t, connMeta, recovered.ConnectionMetadata) + assert.Equal(t, startTime, recovered.StartTime) + assert.Equal(t, lastTime, recovered.EndTime) + assert.Len(t, emitted, 1) + }, + }, + { + name: "App/end already exists", + evts: []apievents.AuditEvent{ + &apievents.AppSessionStart{ + Metadata: apievents.Metadata{Type: events.AppSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + AppMetadata: appMeta, + }, + &apievents.AppSessionEnd{ + Metadata: apievents.Metadata{Type: events.AppSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + AppMetadata: appMeta, + }, + }, + check: alreadyExists( + &apievents.AppSessionEnd{ + Metadata: apievents.Metadata{Type: events.AppSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + AppMetadata: appMeta, + }, + ), + }, + { + name: "App/end recovered", + evts: []apievents.AuditEvent{ + &apievents.AppSessionStart{ + Metadata: apievents.Metadata{Type: events.AppSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + ServerMetadata: serverMeta, + ConnectionMetadata: connMeta, + AppMetadata: appMeta, + }, + &apievents.AppSessionChunk{Metadata: apievents.Metadata{Type: events.AppSessionChunkEvent, Time: lastTime}}, + }, + check: func(t *testing.T, gotEnd apievents.AuditEvent, emitted []apievents.AuditEvent) { + t.Helper() + recovered, ok := gotEnd.(*apievents.AppSessionEnd) + require.True(t, ok) + assert.Equal(t, events.AppSessionEndEvent, recovered.Type) + assert.Equal(t, events.AppSessionEndCode, recovered.Code) + assert.Equal(t, userMeta, recovered.UserMetadata) + assert.Equal(t, sessionMeta, recovered.SessionMetadata) + assert.Equal(t, serverMeta, recovered.ServerMetadata) + assert.Equal(t, connMeta, recovered.ConnectionMetadata) + assert.Equal(t, appMeta, recovered.AppMetadata) + assert.Len(t, emitted, 1) + }, + }, + { + name: "MCP/end already exists", + evts: []apievents.AuditEvent{ + &apievents.MCPSessionStart{ + Metadata: apievents.Metadata{Type: events.MCPSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + AppMetadata: appMeta, + }, + &apievents.MCPSessionEnd{ + Metadata: apievents.Metadata{Type: events.MCPSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + AppMetadata: appMeta, + }, + }, + check: alreadyExists( + &apievents.MCPSessionEnd{ + Metadata: apievents.Metadata{Type: events.MCPSessionEndEvent, Time: lastTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + AppMetadata: appMeta, + }, + ), + }, + { + name: "MCP/end recovered", + evts: []apievents.AuditEvent{ + &apievents.MCPSessionStart{ + Metadata: apievents.Metadata{Type: events.MCPSessionStartEvent, Time: startTime, ClusterName: clusterName}, + UserMetadata: userMeta, + SessionMetadata: sessionMeta, + ServerMetadata: serverMeta, + ConnectionMetadata: connMeta, + AppMetadata: appMeta, + }, + &apievents.SessionPrint{Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: lastTime}}, + }, + check: func(t *testing.T, gotEnd apievents.AuditEvent, emitted []apievents.AuditEvent) { + t.Helper() + recovered, ok := gotEnd.(*apievents.MCPSessionEnd) + require.True(t, ok) + assert.Equal(t, events.MCPSessionEndEvent, recovered.Type) + assert.Equal(t, events.MCPSessionEndCode, recovered.Code) + assert.Equal(t, userMeta, recovered.UserMetadata) + assert.Equal(t, sessionMeta, recovered.SessionMetadata) + assert.Equal(t, serverMeta, recovered.ServerMetadata) + assert.Equal(t, connMeta, recovered.ConnectionMetadata) + assert.Equal(t, appMeta, recovered.AppMetadata) + assert.Len(t, emitted, 1) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + emitter := &eventstest.MockRecorderEmitter{} + cfg := makeConfig(eventstest.NewFakeStreamer(tt.evts, 0), emitter) + gotEnd, err := events.FindOrRecoverSessionEnd(t.Context(), cfg) + if tt.wantErr { + require.Error(t, err) + return + } + require.NoError(t, err) + tt.check(t, gotEnd, emitter.Events()) + }) + } +} diff --git a/lib/events/stream.go b/lib/events/stream.go index b7c5522b497..524f1dcc1db 100644 --- a/lib/events/stream.go +++ b/lib/events/stream.go @@ -138,6 +138,10 @@ type ProtoStreamerConfig struct { SessionSummarizerProvider *summarizer.SessionSummarizerProvider // RecordingMetadataProvider is a provider of the recording metadata service. RecordingMetadataProvider *recordingmetadata.Provider + // OnUploadComplete is called after an upload completes when no session end event + // was observed in the stream. It returns the recovered session end event, if any. + // If nil, no recovery is attempted. + OnUploadComplete func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) } // CheckAndSetDefaults checks and sets streamer defaults @@ -164,16 +168,18 @@ func NewProtoStreamer(cfg ProtoStreamerConfig) (*ProtoStreamer, error) { // Min upload bytes + some overhead to prevent buffer growth (gzip writer is not precise) bufferPool: utils.NewBufferSyncPool(cfg.MinUploadBytes + cfg.MinUploadBytes/3), // MaxProtoMessage size + length of the message record - slicePool: utils.NewSliceSyncPool(constants.MaxProtoMessageSizeBytes + ProtoStreamV1RecordHeaderSize), + slicePool: utils.NewSliceSyncPool(constants.MaxProtoMessageSizeBytes + ProtoStreamV1RecordHeaderSize), + onUploadComplete: cfg.OnUploadComplete, }, nil } // ProtoStreamer creates protobuf-based streams uploaded to the storage // backends, for example S3 or GCS type ProtoStreamer struct { - cfg ProtoStreamerConfig - bufferPool *utils.BufferSyncPool - slicePool *utils.SliceSyncPool + cfg ProtoStreamerConfig + onUploadComplete func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) + bufferPool *utils.BufferSyncPool + slicePool *utils.SliceSyncPool } // CreateAuditStreamForUpload creates audit stream for existing upload, @@ -191,9 +197,18 @@ func (s *ProtoStreamer) CreateAuditStreamForUpload(ctx context.Context, sid sess Encrypter: s.cfg.Encrypter, SessionSummarizerProvider: s.cfg.SessionSummarizerProvider, RecordingMetadataProvider: s.cfg.RecordingMetadataProvider, + OnUploadComplete: s.onUploadComplete, }) } +// SetOnUploadComplete sets a callback to be invoked after an upload completes +// when no session end event was observed in the stream. This allows callers to +// recover or synthesize the session end event from an external source (e.g. +// the audit log). It must be called before any streams are created. +func (s *ProtoStreamer) SetOnUploadComplete(fn func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error)) { + s.onUploadComplete = fn +} + // CreateAuditStream creates audit stream and upload func (s *ProtoStreamer) CreateAuditStream(ctx context.Context, sid session.ID) (apievents.Stream, error) { upload, err := s.cfg.Uploader.CreateUpload(ctx, sid) @@ -223,6 +238,7 @@ func (s *ProtoStreamer) ResumeAuditStream(ctx context.Context, sid session.ID, u Encrypter: s.cfg.Encrypter, SessionSummarizerProvider: s.cfg.SessionSummarizerProvider, RecordingMetadataProvider: s.cfg.RecordingMetadataProvider, + OnUploadComplete: s.onUploadComplete, }) } @@ -263,6 +279,10 @@ type ProtoStreamConfig struct { SessionSummarizerProvider *summarizer.SessionSummarizerProvider // RecordingMetadataProvider is a provider of the recording metadata service. RecordingMetadataProvider *recordingmetadata.Provider + // OnUploadComplete is called after an upload completes when no session end event + // was observed in the stream. It returns the recovered session end event, if any. + // If nil, no recovery is attempted. + OnUploadComplete func(ctx context.Context, sessionID session.ID) (apievents.AuditEvent, error) } // CheckAndSetDefaults checks and sets default values @@ -589,6 +609,9 @@ type sliceWriter struct { // point where the session end event has already been uploaded. If captured, // it will be passed to the summarizer. dbSessionEndEvent *apievents.DatabaseSessionEnd + + // hasSessionEnd indicates if the session end event is present. + hasSessionEnd bool } func (w *sliceWriter) updateCompletedParts(part StreamPart, lastEventIndex int64) { @@ -708,9 +731,17 @@ func (w *sliceWriter) receiveAndUpload() error { case *apievents.OneOf_SessionEnd: w.sshSessionEndEvent = e.SessionEnd w.sessionEndTime = e.SessionEnd.Time + w.hasSessionEnd = true case *apievents.OneOf_DatabaseSessionEnd: w.dbSessionEndEvent = e.DatabaseSessionEnd + w.hasSessionEnd = true + case *apievents.OneOf_WindowsDesktopSessionEnd: + w.hasSessionEnd = true + case *apievents.OneOf_AppSessionEnd: + w.hasSessionEnd = true + case *apievents.OneOf_MCPSessionEnd: + w.hasSessionEnd = true } if w.shouldUploadCurrentSlice() { // this logic blocks the EmitAuditEvent in case if the @@ -839,6 +870,22 @@ func (w *sliceWriter) completeStream() { return } + if !w.hasSessionEnd && w.proto.cfg.OnUploadComplete != nil { + sessionEndEvent, err := w.proto.cfg.OnUploadComplete(w.proto.cancelCtx, w.proto.cfg.Upload.SessionID) + if err != nil { + slog.WarnContext(w.proto.cancelCtx, "Failed to complete upload", "error", err) + return + } + switch o := sessionEndEvent.(type) { + case *apievents.SessionEnd: + w.sshSessionEndEvent = o + w.shouldProcessSession = true + w.sessionEndTime = o.EndTime + case *apievents.DatabaseSessionEnd: + w.dbSessionEndEvent = o + } + } + if w.proto.cfg.RecordingMetadataProvider != nil { recordingMetadata := w.proto.cfg.RecordingMetadataProvider.Service() diff --git a/lib/events/stream_test.go b/lib/events/stream_test.go index 491bca73f8c..147e4e9a5a8 100644 --- a/lib/events/stream_test.go +++ b/lib/events/stream_test.go @@ -39,6 +39,7 @@ import ( apidefaults "github.com/gravitational/teleport/api/defaults" apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/api/utils/keys" + "github.com/gravitational/teleport/lib/auth/recordingmetadata" "github.com/gravitational/teleport/lib/auth/summarizer" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/events/eventstest" @@ -686,3 +687,268 @@ func (m *MockSummarizer) SummarizeWithoutEndEvent(ctx context.Context, sessionID args := m.Called(ctx, sessionID) return args.Error(0) } + +// TestOnUploadComplete_MissingSessionEnd verifies that when a stream is +// completed without a session end event, the OnUploadComplete callback is +// invoked and its returned session end event is passed through for +// summarization and recording metadata processing. +func TestOnUploadComplete_MissingSessionEnd(t *testing.T) { + uploader := eventstest.NewMemoryUploader() + summarizerProvider := &summarizer.SessionSummarizerProvider{} + mockSummarizer := &MockSummarizer{} + summarizerProvider.SetSummarizer(mockSummarizer) + + sid := session.NewID() + + // Build the session end that OnUploadComplete will return. + recoveredEnd := &apievents.SessionEnd{ + Metadata: apievents.Metadata{ + Type: events.SessionEndEvent, + Code: events.SessionEndCode, + }, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + StartTime: time.Now().Add(-time.Minute), + EndTime: time.Now(), + Interactive: true, + } + + called := false + streamer, err := events.NewProtoStreamer(events.ProtoStreamerConfig{ + Uploader: uploader, + SessionSummarizerProvider: summarizerProvider, + }) + require.NoError(t, err) + streamer.SetOnUploadComplete(func(_ context.Context, gotSID session.ID) (apievents.AuditEvent, error) { + called = true + require.Equal(t, sid, gotSID) + return recoveredEnd, nil + }) + + mockSummarizer.On("SummarizeSSH", mock.Anything, mock.MatchedBy(func(e *apievents.SessionEnd) bool { + return e.GetSessionID() == sid.String() + })).Return(nil).Once() + + stream, err := streamer.CreateAuditStream(t.Context(), sid) + require.NoError(t, err) + + preparer, err := events.NewPreparer(events.PreparerConfig{ + SessionID: sid, + Namespace: apidefaults.Namespace, + ClusterName: "cluster", + }) + require.NoError(t, err) + + // Emit a session start but deliberately omit the session end. + start := &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Code: events.SessionStartCode, ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + TerminalSize: "80:25", + } + prepared, err := preparer.PrepareSessionEvent(start) + require.NoError(t, err) + require.NoError(t, stream.RecordEvent(t.Context(), prepared)) + + require.NoError(t, stream.Complete(t.Context())) + + require.True(t, called, "OnUploadComplete must be called when session end is missing") + mockSummarizer.AssertExpectations(t) +} + +// MockRecordingMetadataService is a mock implementation of recordingmetadata.Service. +type MockRecordingMetadataService struct { + mock.Mock +} + +func (m *MockRecordingMetadataService) ProcessSessionRecording(ctx context.Context, sessionID session.ID, sessionType recordingmetadata.SessionType, duration time.Duration) error { + args := m.Called(ctx, sessionID, duration) + return args.Error(0) +} + +// TestRecordingMetadataProcessing verifies that the recording metadata service +// is called with the correct session duration when completing an upload, and +// that sessionEndTime is correctly derived from different event types. +func TestRecordingMetadataProcessing(t *testing.T) { + startTime := time.Date(2024, 1, 15, 10, 0, 0, 0, time.UTC) + + cases := []struct { + name string + buildEvents func(sid session.ID) []apievents.AuditEvent + onUploadComplete func(ctx context.Context, sid session.ID) (apievents.AuditEvent, error) + expectProcess bool + expectedDuration time.Duration + processingError error + }{ + { + name: "sessionEndTime from session end event", + buildEvents: func(sid session.ID) []apievents.AuditEvent { + return []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Code: events.SessionStartCode, Time: startTime, ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + TerminalSize: "80:25", + }, + &apievents.SessionPrint{ + Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: startTime.Add(30 * time.Minute)}, + Data: []byte("hello"), + Bytes: 5, + }, + &apievents.SessionEnd{ + Metadata: apievents.Metadata{Type: events.SessionEndEvent, Code: events.SessionEndCode, Time: startTime.Add(time.Hour), ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + StartTime: startTime, + EndTime: startTime.Add(time.Hour), + Interactive: true, + }, + } + }, + // sessionEndTime is set from SessionEnd.Metadata.Time (not SessionPrint.Time), + // so duration = SessionEnd.Time - SessionStart.Time = 1h. + expectProcess: true, + expectedDuration: time.Hour, + }, + { + name: "sessionEndTime from last print event when no session end", + buildEvents: func(sid session.ID) []apievents.AuditEvent { + return []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Code: events.SessionStartCode, Time: startTime, ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + TerminalSize: "80:25", + }, + &apievents.SessionPrint{ + Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: startTime.Add(30 * time.Minute)}, + Data: []byte("hello"), + Bytes: 5, + }, + } + }, + // No SessionEnd, so sessionEndTime falls back to the last SessionPrint.Time. + expectProcess: true, + expectedDuration: 30 * time.Minute, + }, + { + name: "no processing when session start event is missing", + buildEvents: func(sid session.ID) []apievents.AuditEvent { + return []apievents.AuditEvent{ + &apievents.SessionPrint{ + Metadata: apievents.Metadata{Type: events.SessionPrintEvent, Time: startTime.Add(5 * time.Minute)}, + Data: []byte("hello"), + Bytes: 5, + }, + } + }, + // shouldProcessSession is never set without a SessionStart. + expectProcess: false, + }, + { + name: "no processing when session end time is zero", + buildEvents: func(sid session.ID) []apievents.AuditEvent { + return []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Code: events.SessionStartCode, Time: startTime, ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + TerminalSize: "80:25", + }, + } + }, + // shouldProcessSession is true, but sessionEndTime is zero (no prints or end event). + expectProcess: false, + }, + { + name: "sessionEndTime from OnUploadComplete recovered SessionEnd", + buildEvents: func(sid session.ID) []apievents.AuditEvent { + return []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Code: events.SessionStartCode, Time: startTime, ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + TerminalSize: "80:25", + }, + } + }, + onUploadComplete: func(_ context.Context, gotSID session.ID) (apievents.AuditEvent, error) { + return &apievents.SessionEnd{ + Metadata: apievents.Metadata{Type: events.SessionEndEvent, Code: events.SessionEndCode}, + SessionMetadata: apievents.SessionMetadata{SessionID: gotSID.String()}, + StartTime: startTime, + EndTime: startTime.Add(45 * time.Minute), + Interactive: true, + }, nil + }, + // sessionEndTime is set from the recovered SessionEnd.EndTime. + expectProcess: true, + expectedDuration: 45 * time.Minute, + }, + { + name: "processing error does not cause panic", + buildEvents: func(sid session.ID) []apievents.AuditEvent { + return []apievents.AuditEvent{ + &apievents.SessionStart{ + Metadata: apievents.Metadata{Type: events.SessionStartEvent, Code: events.SessionStartCode, Time: startTime, ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + TerminalSize: "80:25", + }, + &apievents.SessionEnd{ + Metadata: apievents.Metadata{Type: events.SessionEndEvent, Code: events.SessionEndCode, Time: startTime.Add(time.Hour), ClusterName: "cluster"}, + SessionMetadata: apievents.SessionMetadata{SessionID: sid.String()}, + StartTime: startTime, + EndTime: startTime.Add(time.Hour), + Interactive: true, + }, + } + }, + expectProcess: true, + expectedDuration: time.Hour, + processingError: errors.New("processing error"), + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + summarizerProvider := &summarizer.SessionSummarizerProvider{} + metadataProvider := &recordingmetadata.Provider{} + uploader := eventstest.NewMemoryUploader() + + streamer, err := events.NewProtoStreamer(events.ProtoStreamerConfig{ + Uploader: uploader, + SessionSummarizerProvider: summarizerProvider, + RecordingMetadataProvider: metadataProvider, + }) + require.NoError(t, err) + + sid := session.NewID() + mockMetadata := &MockRecordingMetadataService{} + metadataProvider.SetService(mockMetadata) + + if tc.expectProcess { + mockMetadata. + On("ProcessSessionRecording", mock.Anything, sid, tc.expectedDuration). + Return(tc.processingError). + Once() + } + + if tc.onUploadComplete != nil { + streamer.SetOnUploadComplete(tc.onUploadComplete) + } + + stream, err := streamer.CreateAuditStream(t.Context(), sid) + require.NoError(t, err) + + preparer, err := events.NewPreparer(events.PreparerConfig{ + SessionID: sid, + Namespace: apidefaults.Namespace, + ClusterName: "cluster", + }) + require.NoError(t, err) + + for _, evt := range tc.buildEvents(sid) { + prepared, err := preparer.PrepareSessionEvent(evt) + require.NoError(t, err) + require.NoError(t, stream.RecordEvent(t.Context(), prepared)) + } + + require.NoError(t, stream.Complete(t.Context())) + + mockMetadata.AssertExpectations(t) + }) + } +} diff --git a/lib/service/service.go b/lib/service/service.go index 45d40d612a3..18357e85e97 100644 --- a/lib/service/service.go +++ b/lib/service/service.go @@ -2304,7 +2304,7 @@ func (process *TeleportProcess) initAuthService() error { clusterConfig = recordingEncryptionManager var emitter apievents.Emitter - var streamer events.Streamer + var streamer events.StreamerWithCallback var uploadHandler events.MultipartHandler var externalAuditStorage *externalauditstorage.Configurator encryptedIO, err := recordingencryption.NewEncryptedIO(clusterConfig, recordingEncryptionManager) @@ -2317,6 +2317,7 @@ func (process *TeleportProcess) initAuthService() error { // create the audit log, which will be consuming (and recording) all events // and recording all sessions. + var localLog *events.AuditLog if cfg.Auth.NoAudit { // this is for teleconsole process.auditLog = events.NewDiscardAuditLog() @@ -2384,7 +2385,7 @@ func (process *TeleportProcess) initAuthService() error { if err != nil { return trace.Wrap(err) } - localLog, err := events.NewAuditLog(auditServiceConfig) + localLog, err = events.NewAuditLog(auditServiceConfig) if err != nil { return trace.Wrap(err) } @@ -2530,7 +2531,12 @@ func (process *TeleportProcess) initAuthService() error { } authServer.EncryptedIO = encryptedIO - + if streamer != nil { + streamer.SetOnUploadComplete(authServer.OnUploadComplete) + } + if localLog != nil { + localLog.SetOnUploadComplete(authServer.OnUploadComplete) + } lockWatcher, err := services.NewLockWatcher(process.ExitContext(), services.LockWatcherConfig{ ResourceWatcherConfig: services.ResourceWatcherConfig{ Component: teleport.ComponentAuth, @@ -2626,6 +2632,7 @@ func (process *TeleportProcess) initAuthService() error { ServerID: hostUUID, SessionSummarizerProvider: sessionSummarizerProvider, RecordingMetadataProvider: recordingMetadataProvider, + EnsureSessionEndEvent: true, }) if err != nil { return trace.Wrap(err, "starting upload completer") diff --git a/lib/services/role.go b/lib/services/role.go index 35b54a45bfc..ca1495cfd24 100644 --- a/lib/services/role.go +++ b/lib/services/role.go @@ -940,6 +940,12 @@ func RoleSetFromSpec(name string, spec types.RoleSpecV6) (RoleSet, error) { return NewRoleSet(role), nil } +// WO is a shortcut that returns create and update verbs, granting the ability +// to emit/write resources but not list, read, or delete them. +func WO() []string { + return []string{types.VerbCreate, types.VerbUpdate} +} + // RW is a shortcut that returns all CRUD verbs. func RW() []string { return []string{types.VerbList, types.VerbCreate, types.VerbRead, types.VerbUpdate, types.VerbDelete}