Prevent teleport built-in roles from accessing audit log and sessions (#66079)

Teleport RBAC policies give built-in role certificates (proxy,
discovery, okta, node, kube...) access to the session recordings
API and audit log.

For the session recordings API, the policy enforcement completely
skips the RBAC logic and if the certificate was from a built-in role,
the client verification logic would skip RBAC entirely.

```go
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)
	}

```

From the built-in roles, only node had read access to
`types.KindSession`, but given the check above all of them had access,
including cases like `okta` where no session recording exists.
This means a user with access to a token or an agent certificate can
read, play and view every single session recording.

Historically, this bypass existed for one particular reason: when an
upload is abandoned without proper termination - this happens when the
agent is restarted while interactive sessions are alive, the agent
places the session in a temporary folder. This folder is constantly
monitored by a routine called `UploadCompleter`. This routine runs on
every teleport process and is responsible for deciding when a session
was abandoned - no active session tracker and the last written data was
24h ago and mark it completed, i.e. move it to another directory, so that
the file uploader could move it to auth server for long-term storage.

Since the session was not correctly terminated it could be the case it
misses the `session.end` or respective end event. This is not
particularly important for the recording itself since it's playable as is,
but it's critical for the session recordings list.

When a user goes to Session Recordings page or does `tctl recordings ls`,
Teleport issues a read request from the audit log to search for
`session.end`, `windows.desktop.session.end`, `db.session.end` and
`app.session.end`. If the audit log doesn't contain the events but the
agent successfully uploaded the recording, the recording won't be
included in the list and will be inaccessible from `tctl` or `tsh`.

To overcome this problem,
https://github.com/gravitational/teleport/pull/14521 introduced logic to
reconstruct the session end event from the recording itself. The logic
assumes (partially incorrect) that if the session recording holds the
session end event, the same event was already present in the audit log.
After the `UploadCompleter` successfully marked the upload as completed,
it scheduled a goroutine to stream the session and see if the session end
event exists. The goroutine had to run 2m after because it was considered
a good time for the file uploader to send the session to the auth server
and auth made it available in the long term storage. To stream it, the
agent contacted Auth Server through the same API users use to play sessions.

Because of an improper filtering and logic, any agent had access to any session.
Even without knowing the session id, given that the agent also had access to read
the audit log, it was able to search for the session end events and later access
them using the `StreamSessionEvents` API.
The main reason why it was implemented like this was because from the moment the
`UploadCompleter` marked the session as completed, auth server had no way to distinguish
between a complete session vs an incomplete session. So the logic was placed into the last
piece that still retained that information.

This PR fixes both problems. It removes RW access to audit logs,
lowering it to WO - agents can write it but can't read it, and removed
any access to StreamSessionEvents.

In order to remove the `StreamSessionEvents` logic being accessible by
the agents, this PR moves the same logic to auth server. There are 3
places where this logic must exist:

- Auth upload completer: If the cluster is operating with any sync mode
  instead of async, Auth is the receiving point the the streams. If
  something fails, the recording will live in auth server and not the
  agent.
- Upload API in Auth: When clusters operate in async mode, agents'
  fileuploader uses the auth grpc API to upload the events. Since auth
  receives event by event, we can easily analyse if the session already
  has a session end event and fill it if necessary
- Encrypted recordings upload: If the cluster operates with encrypted
  recordings enabled, agents upload the data to the encrypted recordings
  gRPC service. Before this fix, if the session was incomplete we didn't
  even try to build the session end event and the session was hidden.

This PR exposes the completer logic in auth server and each one of the 3
cases mentioned call that auth's server function for consistency and
predictability.

Fixes https://github.com/gravitational/teleport-private/issues/2363
This PR fixes https://github.com/gravitational/teleport/issues/60861
accidentally.
This commit is contained in:
Tiago Silva
2026-04-24 16:41:58 +00:00
committed by GitHub
parent f041523934
commit 0cc28f8a79
20 changed files with 1382 additions and 321 deletions
+67 -30
View File
@@ -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 {
+28 -2
View File
@@ -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)
+9 -9
View File
@@ -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)
+93 -42
View File
@@ -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,
+1 -1
View File
@@ -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)
@@ -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
}
@@ -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{
+9 -11
View File
@@ -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()),
+12
View File
@@ -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
+14 -1
View File
@@ -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()
+8
View File
@@ -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()
+65 -137
View File
@@ -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
}
+56 -51
View File
@@ -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())
})
}
}
+4
View File
@@ -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{}
+207
View File
@@ -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
}
+376
View File
@@ -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())
})
}
}
+51 -4
View File
@@ -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()
+266
View File
@@ -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)
})
}
}
+10 -3
View File
@@ -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")
+6
View File
@@ -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}