diff --git a/lib/bpf/bpf.go b/lib/bpf/bpf.go index 7fd4bb5c3a1..7fb341d2c37 100644 --- a/lib/bpf/bpf.go +++ b/lib/bpf/bpf.go @@ -38,6 +38,7 @@ import ( apievents "github.com/gravitational/teleport/api/types/events" controlgroup "github.com/gravitational/teleport/lib/cgroup" "github.com/gravitational/teleport/lib/events" + "github.com/gravitational/teleport/lib/srv" "github.com/gravitational/trace" "github.com/gravitational/ttlmap" @@ -53,24 +54,24 @@ const ArgsCacheSize = 1024 // SessionWatch is a map of cgroup IDs that the BPF service is watching and // emitting events for. type SessionWatch struct { - watch map[uint64]*SessionContext + watch map[uint64]*srv.ServerContext mu sync.Mutex } func NewSessionWatch() SessionWatch { return SessionWatch{ - watch: make(map[uint64]*SessionContext), + watch: make(map[uint64]*srv.ServerContext), } } -func (w *SessionWatch) Get(cgoupID uint64) (ctx *SessionContext, ok bool) { +func (w *SessionWatch) Get(cgoupID uint64) (ctx *srv.ServerContext, ok bool) { w.mu.Lock() defer w.mu.Unlock() ctx, ok = w.watch[cgoupID] return } -func (w *SessionWatch) Add(cgroupID uint64, ctx *SessionContext) { +func (w *SessionWatch) Add(cgroupID uint64, ctx *srv.ServerContext) { w.mu.Lock() defer w.mu.Unlock() @@ -210,13 +211,14 @@ func (s *Service) Close() error { // OpenSession will place a process within a cgroup and being monitoring all // events from that cgroup and emitting the results to the audit log. -func (s *Service) OpenSession(ctx *SessionContext) (uint64, error) { - err := s.cgroup.Create(ctx.SessionID) +func (s *Service) OpenSession(ctx *srv.ServerContext) (uint64, error) { + sessionID := ctx.SessionID() + err := s.cgroup.Create(sessionID.String()) if err != nil { return 0, trace.Wrap(err) } - cgroupID, err := s.cgroup.ID(ctx.SessionID) + cgroupID, err := s.cgroup.ID(sessionID.String()) if err != nil { return 0, trace.Wrap(err) } @@ -225,7 +227,7 @@ func (s *Service) OpenSession(ctx *SessionContext) (uint64, error) { s.watch.Add(cgroupID, ctx) // Place requested PID into cgroup. - err = s.cgroup.Place(ctx.SessionID, ctx.PID) + err = s.cgroup.Place(sessionID.String(), ctx.GetPID()) if err != nil { return 0, trace.Wrap(err) } @@ -235,8 +237,9 @@ func (s *Service) OpenSession(ctx *SessionContext) (uint64, error) { // CloseSession will stop monitoring events from a particular cgroup and // remove the cgroup. -func (s *Service) CloseSession(ctx *SessionContext) error { - cgroupID, err := s.cgroup.ID(ctx.SessionID) +func (s *Service) CloseSession(ctx *srv.ServerContext) error { + sessionID := ctx.SessionID() + cgroupID, err := s.cgroup.ID(sessionID.String()) if err != nil { return trace.Wrap(err) } @@ -246,7 +249,7 @@ func (s *Service) CloseSession(ctx *SessionContext) error { // Move all PIDs to the root cgroup and remove the cgroup created for this // session. - err = s.cgroup.Remove(ctx.SessionID) + err = s.cgroup.Remove(sessionID.String()) if err != nil { return trace.Wrap(err) } @@ -294,7 +297,7 @@ func (s *Service) emitCommandEvent(eventBytes []byte) { } // If the command event is not being monitored, don't process it. - _, ok = ctx.Events[constants.EnhancedRecordingCommand] + _, ok = ctx.Identity.AccessChecker.EnhancedRecordingSet()[constants.EnhancedRecordingCommand] if !ok { return } @@ -326,21 +329,22 @@ func (s *Service) emitCommandEvent(eventBytes []byte) { argv := args.([]string) // Emit "command" event. + sessionID := ctx.SessionID() sessionCommandEvent := &apievents.SessionCommand{ Metadata: apievents.Metadata{ Type: events.SessionCommandEvent, Code: events.SessionCommandCode, }, ServerMetadata: apievents.ServerMetadata{ - ServerID: ctx.ServerID, - ServerNamespace: ctx.Namespace, + ServerID: ctx.GetServer().HostUUID(), + ServerNamespace: ctx.GetServer().GetNamespace(), }, SessionMetadata: apievents.SessionMetadata{ - SessionID: ctx.SessionID, + SessionID: sessionID.String(), }, UserMetadata: apievents.UserMetadata{ - User: ctx.User, - Login: ctx.Login, + User: ctx.Identity.TeleportUser, + Login: ctx.Identity.Login, }, BPFMetadata: apievents.BPFMetadata{ CgroupID: event.CgroupID, @@ -352,7 +356,7 @@ func (s *Service) emitCommandEvent(eventBytes []byte) { Path: argv[0], Argv: argv[1:], } - if err := ctx.Emitter.EmitAuditEvent(ctx.Context, sessionCommandEvent); err != nil { + if err := ctx.StreamWriter().EmitAuditEvent(ctx.Context, sessionCommandEvent); err != nil { log.WithError(err).Warn("Failed to emit command event.") } @@ -378,26 +382,27 @@ func (s *Service) emitDiskEvent(eventBytes []byte) { } // If the network event is not being monitored, don't process it. - _, ok = ctx.Events[constants.EnhancedRecordingDisk] + _, ok = ctx.Identity.AccessChecker.EnhancedRecordingSet()[constants.EnhancedRecordingDisk] if !ok { return } + sessionID := ctx.SessionID() sessionDiskEvent := &apievents.SessionDisk{ Metadata: apievents.Metadata{ Type: events.SessionDiskEvent, Code: events.SessionDiskCode, }, ServerMetadata: apievents.ServerMetadata{ - ServerID: ctx.ServerID, - ServerNamespace: ctx.Namespace, + ServerID: ctx.GetServer().HostUUID(), + ServerNamespace: ctx.GetServer().GetNamespace(), }, SessionMetadata: apievents.SessionMetadata{ - SessionID: ctx.SessionID, + SessionID: sessionID.String(), }, UserMetadata: apievents.UserMetadata{ - User: ctx.User, - Login: ctx.Login, + User: ctx.Identity.TeleportUser, + Login: ctx.Identity.Login, }, BPFMetadata: apievents.BPFMetadata{ CgroupID: event.CgroupID, @@ -409,7 +414,7 @@ func (s *Service) emitDiskEvent(eventBytes []byte) { ReturnCode: event.ReturnCode, } // Logs can be DoS by event failures here - _ = ctx.Emitter.EmitAuditEvent(ctx.Context, sessionDiskEvent) + _ = ctx.StreamWriter().EmitAuditEvent(ctx.Context, sessionDiskEvent) } // emit4NetworkEvent will parse and emit IPv4 events to the Audit Log. @@ -429,7 +434,7 @@ func (s *Service) emit4NetworkEvent(eventBytes []byte) { } // If the network event is not being monitored, don't process it. - _, ok = ctx.Events[constants.EnhancedRecordingNetwork] + _, ok = ctx.Identity.AccessChecker.EnhancedRecordingSet()[constants.EnhancedRecordingNetwork] if !ok { return } @@ -444,21 +449,22 @@ func (s *Service) emit4NetworkEvent(eventBytes []byte) { binary.LittleEndian.PutUint32(dst, event.DstAddr) dstAddr := net.IP(dst) + sessionID := ctx.SessionID() sessionNetworkEvent := &apievents.SessionNetwork{ Metadata: apievents.Metadata{ Type: events.SessionNetworkEvent, Code: events.SessionNetworkCode, }, ServerMetadata: apievents.ServerMetadata{ - ServerID: ctx.ServerID, - ServerNamespace: ctx.Namespace, + ServerID: ctx.GetServer().HostUUID(), + ServerNamespace: ctx.GetServer().GetNamespace(), }, SessionMetadata: apievents.SessionMetadata{ - SessionID: ctx.SessionID, + SessionID: sessionID.String(), }, UserMetadata: apievents.UserMetadata{ - User: ctx.User, - Login: ctx.Login, + User: ctx.Identity.TeleportUser, + Login: ctx.Identity.Login, }, BPFMetadata: apievents.BPFMetadata{ CgroupID: event.CgroupID, @@ -470,7 +476,7 @@ func (s *Service) emit4NetworkEvent(eventBytes []byte) { SrcAddr: srcAddr.String(), TCPVersion: 4, } - if err := ctx.Emitter.EmitAuditEvent(ctx.Context, sessionNetworkEvent); err != nil { + if err := ctx.StreamWriter().EmitAuditEvent(ctx.Context, sessionNetworkEvent); err != nil { log.WithError(err).Warn("Failed to emit network event.") } } @@ -492,7 +498,7 @@ func (s *Service) emit6NetworkEvent(eventBytes []byte) { } // If the network event is not being monitored, don't process it. - _, ok = ctx.Events[constants.EnhancedRecordingNetwork] + _, ok = ctx.Identity.AccessChecker.EnhancedRecordingSet()[constants.EnhancedRecordingNetwork] if !ok { return } @@ -513,21 +519,22 @@ func (s *Service) emit6NetworkEvent(eventBytes []byte) { binary.LittleEndian.PutUint32(dst[12:], event.DstAddr[3]) dstAddr := net.IP(dst) + sessionID := ctx.SessionID() sessionNetworkEvent := &apievents.SessionNetwork{ Metadata: apievents.Metadata{ Type: events.SessionNetworkEvent, Code: events.SessionNetworkCode, }, ServerMetadata: apievents.ServerMetadata{ - ServerID: ctx.ServerID, - ServerNamespace: ctx.Namespace, + ServerID: ctx.GetServer().HostUUID(), + ServerNamespace: ctx.GetServer().GetNamespace(), }, SessionMetadata: apievents.SessionMetadata{ - SessionID: ctx.SessionID, + SessionID: sessionID.String(), }, UserMetadata: apievents.UserMetadata{ - User: ctx.User, - Login: ctx.Login, + User: ctx.Identity.TeleportUser, + Login: ctx.Identity.Login, }, BPFMetadata: apievents.BPFMetadata{ CgroupID: event.CgroupID, @@ -539,7 +546,7 @@ func (s *Service) emit6NetworkEvent(eventBytes []byte) { SrcAddr: srcAddr.String(), TCPVersion: 6, } - if err := ctx.Emitter.EmitAuditEvent(ctx.Context, sessionNetworkEvent); err != nil { + if err := ctx.StreamWriter().EmitAuditEvent(ctx.Context, sessionNetworkEvent); err != nil { log.WithError(err).Warn("Failed to emit network event.") } } diff --git a/lib/bpf/bpf_test.go b/lib/bpf/bpf_test.go index 00125d5738f..149bc48da41 100644 --- a/lib/bpf/bpf_test.go +++ b/lib/bpf/bpf_test.go @@ -34,17 +34,27 @@ import ( "unsafe" "github.com/gravitational/teleport/api/constants" - apidefaults "github.com/gravitational/teleport/api/defaults" + "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/events/eventstest" + "github.com/gravitational/teleport/lib/services" + "github.com/gravitational/teleport/lib/srv" "github.com/aquasecurity/libbpfgo" - "github.com/google/uuid" "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/testutil" "github.com/stretchr/testify/require" ) +type terminal struct { + srv.Terminal + pid int +} + +func (t terminal) PID() int { + return t.pid +} + func TestRootWatch(t *testing.T) { // TODO(jakule): Find a way to run this test in CI. Disable for now to not block all BPF tests. t.Skip("this test always fails when running inside a CGroup/Docker") @@ -79,23 +89,23 @@ func TestRootWatch(t *testing.T) { err = cmd.Start() require.NoError(t, err) - // Create a monitoring session for init. The events we execute should not - // have PID 1, so nothing should be captured in the Audit Log. - cgroupID, err := service.OpenSession(&SessionContext{ - Namespace: apidefaults.Namespace, - SessionID: uuid.New().String(), - ServerID: uuid.New().String(), - Login: "foo", - User: "foo@example.com", - PID: cmd.Process.Pid, - Emitter: emitter, - Events: map[string]bool{ - constants.EnhancedRecordingCommand: true, - constants.EnhancedRecordingDisk: true, - constants.EnhancedRecordingNetwork: true, + server := srv.NewMockServer(t) + server.MockEmitter = emitter + + role, err := types.NewRole("bpf", types.RoleSpecV5{ + Options: types.RoleOptions{ + BPF: []string{constants.EnhancedRecordingCommand, constants.EnhancedRecordingDisk, constants.EnhancedRecordingNetwork}, }, }) require.NoError(t, err) + + srvctx := srv.NewTestServerContext(t, server, services.NewRoleSet(role)) + srvctx.SetTerm(terminal{pid: os.Getpid()}) + + // Create a monitoring session for init. The events we execute should not + // have PID 1, so nothing should be captured in the Audit Log. + cgroupID, err := service.OpenSession(srvctx) + require.NoError(t, err) require.Greater(t, cgroupID, 0) // Find "ls" binary. diff --git a/lib/bpf/common.go b/lib/bpf/common.go index 6516eceabfb..1439aa133d3 100644 --- a/lib/bpf/common.go +++ b/lib/bpf/common.go @@ -22,11 +22,9 @@ package bpf import "C" import ( - "context" - "github.com/gravitational/teleport/api/constants" - apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/defaults" + "github.com/gravitational/teleport/lib/srv" "github.com/gravitational/teleport/lib/utils" "github.com/gravitational/trace" @@ -38,50 +36,15 @@ import ( type BPF interface { // OpenSession will start monitoring all events within a session and // emitting them to the Audit Log. - OpenSession(ctx *SessionContext) (uint64, error) + OpenSession(ctx *srv.ServerContext) (uint64, error) // CloseSession will stop monitoring events for a particular session. - CloseSession(ctx *SessionContext) error + CloseSession(ctx *srv.ServerContext) error // Close will stop any running BPF programs. Close() error } -// SessionContext contains all the information needed to track and emit -// events for a particular session. Most of this information is already within -// srv.ServerContext, unfortunately due to circular imports with lib/srv and -// lib/bpf, part of that structure is reproduced in SessionContext. -type SessionContext struct { - // Context is a cancel context, scoped to a server, and not a session. - Context context.Context - - // Namespace is the namespace within which this session occurs. - Namespace string - - // SessionID is the UUID of the given session. - SessionID string - - // ServerID is the UUID of the server this session is executing on. - ServerID string - - // Login is the Unix login for this session. - Login string - - // User is the Teleport user. - User string - - // PID is the process ID of Teleport when it re-executes itself. This is - // used by Teleport to find itself by cgroup. - PID int - - // Emitter is used to record events for a particular session - Emitter apievents.Emitter - - // Events is the set of events (command, disk, or network) to record for - // this session. - Events map[string]bool -} - // Config holds configuration for the BPF service. type Config struct { // Enabled is if this service will try and install BPF programs on this system. @@ -122,8 +85,7 @@ func (c *Config) CheckAndSetDefaults() error { } // NOP is used on either non-Linux systems or when BPF support is not enabled. -type NOP struct { -} +type NOP struct{} // Close closes the NOP service. Note this function does nothing. func (s *NOP) Close() error { @@ -131,12 +93,12 @@ func (s *NOP) Close() error { } // OpenSession opens a NOP session. Note this function does nothing. -func (s *NOP) OpenSession(_ *SessionContext) (uint64, error) { +func (s *NOP) OpenSession(_ *srv.ServerContext) (uint64, error) { return 0, nil } // CloseSession closes a NOP session. Note this function does nothing. -func (s *NOP) CloseSession(_ *SessionContext) error { +func (s *NOP) CloseSession(_ *srv.ServerContext) error { return nil } diff --git a/lib/bpf/helper.go b/lib/bpf/helper.go index 61dabe04988..1e380cf65e2 100644 --- a/lib/bpf/helper.go +++ b/lib/bpf/helper.go @@ -218,7 +218,7 @@ func (c *Counter) Close() { } func (c *Counter) loop() { - for _ = range c.doorbellCh { + for range c.doorbellCh { var key int32 = 0 cntBytes, err := c.arr.GetValue(unsafe.Pointer(&key)) if err != nil { diff --git a/lib/restrictedsession/audit.go b/lib/restrictedsession/audit.go index 5ad1e7756a3..8334b682c78 100644 --- a/lib/restrictedsession/audit.go +++ b/lib/restrictedsession/audit.go @@ -29,6 +29,8 @@ import ( "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/bpf" api "github.com/gravitational/teleport/lib/events" + "github.com/gravitational/teleport/lib/srv" + "github.com/gravitational/trace" ) @@ -67,22 +69,24 @@ type auditEventBlockedIPv6 struct { // newNetworkAuditEvent creates events.SessionNetwork, filling in common fields // from the SessionContext -func newNetworkAuditEvent(ctx *bpf.SessionContext, hdr *auditEventHeader) events.SessionNetwork { +func newNetworkAuditEvent(ctx *srv.ServerContext, hdr *auditEventHeader) events.SessionNetwork { + + sessionID := ctx.SessionID() return events.SessionNetwork{ Metadata: events.Metadata{ Type: api.SessionNetworkEvent, Code: api.SessionNetworkCode, }, ServerMetadata: events.ServerMetadata{ - ServerID: ctx.ServerID, - ServerNamespace: ctx.Namespace, + ServerID: ctx.GetServer().HostUUID(), + ServerNamespace: ctx.GetServer().GetNamespace(), }, SessionMetadata: events.SessionMetadata{ - SessionID: ctx.SessionID, + SessionID: sessionID.String(), }, UserMetadata: events.UserMetadata{ - User: ctx.User, - Login: ctx.Login, + User: ctx.Identity.TeleportUser, + Login: ctx.Identity.Login, }, BPFMetadata: events.BPFMetadata{ CgroupID: hdr.CGroupID, @@ -117,7 +121,7 @@ func ip6String(ip net.IP) string { } // parseAuditEvent parses the body of the audit event -func parseAuditEvent(buf *bytes.Buffer, hdr *auditEventHeader, ctx *bpf.SessionContext) (events.AuditEvent, error) { +func parseAuditEvent(buf *bytes.Buffer, hdr *auditEventHeader, ctx *srv.ServerContext) (events.AuditEvent, error) { switch hdr.EventType { case BlockedIP4: var body auditEventBlockedIPv4 @@ -143,8 +147,8 @@ func parseAuditEvent(buf *bytes.Buffer, hdr *auditEventHeader, ctx *bpf.SessionC event := newNetworkAuditEvent(ctx, hdr) event.DstPort = int32(body.DstPort) - event.DstAddr = ip6String(net.IP(body.DstIP[:])) - event.SrcAddr = ip6String(net.IP(body.SrcIP[:])) + event.DstAddr = ip6String(body.DstIP[:]) + event.SrcAddr = ip6String(body.SrcIP[:]) event.TCPVersion = 6 event.Operation = events.SessionNetwork_NetworkOperation(body.Op) event.Action = events.EventAction_DENIED diff --git a/lib/restrictedsession/manager.go b/lib/restrictedsession/manager.go index 94a1b21bf81..012ac307fd4 100644 --- a/lib/restrictedsession/manager.go +++ b/lib/restrictedsession/manager.go @@ -18,8 +18,8 @@ package restrictedsession import ( "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/bpf" "github.com/gravitational/teleport/lib/services" + "github.com/gravitational/teleport/lib/srv" ) // RestrictionsWatcherClient is used by changeset to fetch a list @@ -32,20 +32,20 @@ type RestrictionsWatcherClient interface { // Manager starts and stop enforcing restrictions for a given session. type Manager interface { // OpenSession starts enforcing restrictions for a cgroup with cgroupID - OpenSession(ctx *bpf.SessionContext, cgroupID uint64) + OpenSession(ctx *srv.ServerContext, cgroupID uint64) // CloseSession stops enforcing restrictions for a cgroup with cgroupID - CloseSession(ctx *bpf.SessionContext, cgroupID uint64) + CloseSession(ctx *srv.ServerContext, cgroupID uint64) // Close stops the manager, cleaning up any resources Close() } -// Stubbed out Manager interface for cases where the real thing is not used. +// NOP is a stubbed implementation of Manager for cases where the real thing is not used. type NOP struct{} -func (NOP) OpenSession(ctx *bpf.SessionContext, cgroupID uint64) { +func (NOP) OpenSession(ctx *srv.ServerContext, cgroupID uint64) { } -func (NOP) CloseSession(ctx *bpf.SessionContext, cgroupID uint64) { +func (NOP) CloseSession(ctx *srv.ServerContext, cgroupID uint64) { } func (NOP) UpdateNetworkRestrictions(r *NetworkRestrictions) error { diff --git a/lib/restrictedsession/restricted.go b/lib/restrictedsession/restricted.go index c406ca024a0..29164a85712 100644 --- a/lib/restrictedsession/restricted.go +++ b/lib/restrictedsession/restricted.go @@ -27,13 +27,14 @@ import ( "sync" "unsafe" - "github.com/gravitational/teleport" - "github.com/gravitational/teleport/lib/bpf" + "github.com/aquasecurity/libbpfgo" "github.com/gravitational/trace" "github.com/prometheus/client_golang/prometheus" - - "github.com/aquasecurity/libbpfgo" "github.com/sirupsen/logrus" + + "github.com/gravitational/teleport" + "github.com/gravitational/teleport/lib/bpf" + "github.com/gravitational/teleport/lib/srv" ) var log = logrus.WithFields(logrus.Fields{ @@ -163,7 +164,7 @@ func (m *sessionMgr) Close() { // OpenSession inserts the cgroupID into the BPF hash map to enable // enforcement by the kernel -func (m *sessionMgr) OpenSession(ctx *bpf.SessionContext, cgroupID uint64) { +func (m *sessionMgr) OpenSession(ctx *srv.ServerContext, cgroupID uint64) { m.watch.Add(cgroupID, ctx) key := make([]byte, 8) @@ -176,7 +177,7 @@ func (m *sessionMgr) OpenSession(ctx *bpf.SessionContext, cgroupID uint64) { // CloseSession removes the cgroupID from the BPF hash map to enable // enforcement by the kernel -func (m *sessionMgr) CloseSession(ctx *bpf.SessionContext, cgroupID uint64) { +func (m *sessionMgr) CloseSession(ctx *srv.ServerContext, cgroupID uint64) { key := make([]byte, 8) binary.LittleEndian.PutUint64(key, cgroupID) @@ -291,7 +292,7 @@ func (l *auditEventLoop) loop() { continue } - if err = ctx.Emitter.EmitAuditEvent(ctx.Context, event); err != nil { + if err = ctx.StreamWriter().EmitAuditEvent(ctx.Context, event); err != nil { log.WithError(err).Warn("Failed to emit network event.") } } diff --git a/lib/restrictedsession/restricted_test.go b/lib/restrictedsession/restricted_test.go index a77afdb7e20..9b0d59fd110 100644 --- a/lib/restrictedsession/restricted_test.go +++ b/lib/restrictedsession/restricted_test.go @@ -30,17 +30,16 @@ import ( "testing" "time" - apidefaults "github.com/gravitational/teleport/api/defaults" api "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/bpf" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/events/eventstest" "github.com/gravitational/teleport/lib/services" + "github.com/gravitational/teleport/lib/srv" "github.com/gravitational/teleport/lib/utils" go_cmp "github.com/google/go-cmp/cmp" - "github.com/google/uuid" "github.com/stretchr/testify/require" ) @@ -64,7 +63,7 @@ const ( var ( testRanges = []blockedRange{ - blockedRange{ + { ver: 4, allow: "39.156.69.70/28", deny: "39.156.69.71", @@ -77,7 +76,7 @@ var ( "72.156.69.80": denied, }, }, - blockedRange{ + { ver: 4, allow: "77.88.55.88", probe: map[string]blockAction{ @@ -87,7 +86,7 @@ var ( "67.88.55.86": denied, }, }, - blockedRange{ + { ver: 6, allow: "39.156.68.48/28", deny: "39.156.68.48/31", @@ -101,7 +100,7 @@ var ( "::ffff:72.156.68.80": denied, }, }, - blockedRange{ + { ver: 6, allow: "fc80::/64", deny: "fc80::10/124", @@ -114,7 +113,7 @@ var ( "fc60:0:0:1::": denied, }, }, - blockedRange{ + { ver: 6, allow: "2607:f8b0:4005:80a::200e", probe: map[string]blockAction{ @@ -130,7 +129,7 @@ var ( type bpfContext struct { cgroupDir string cgroupID uint64 - ctx *bpf.SessionContext + ctx *srv.ServerContext enhancedRecorder bpf.BPF restrictedMgr Manager srcAddrs map[int]string @@ -140,6 +139,15 @@ type bpfContext struct { expectedAuditEvents []apievents.AuditEvent } +type terminal struct { + srv.Terminal + pid int +} + +func (t terminal) PID() int { + return t.pid +} + func setupBPFContext(t *testing.T) *bpfContext { tt := bpfContext{} t.Cleanup(func() { tt.Close(t) }) @@ -168,25 +176,24 @@ func setupBPFContext(t *testing.T) *bpfContext { }) require.NoError(t, err) - // Create the SessionContext used by both enhanced recording and us (restricted session) - tt.ctx = &bpf.SessionContext{ - Namespace: apidefaults.Namespace, - SessionID: uuid.New().String(), - ServerID: uuid.New().String(), - Login: "foo", - User: "foo@example.com", - PID: os.Getpid(), - Emitter: &tt.emitter, - Events: map[string]bool{}, - } + server := srv.NewMockServer(t) + server.MockEmitter = &tt.emitter + + role, err := api.NewRole("restricted", api.RoleSpecV5{}) + require.NoError(t, err) + + srvctx := srv.NewTestServerContext(t, server, services.NewRoleSet(role)) + srvctx.SetTerm(terminal{pid: os.Getpid()}) + + tt.ctx = srvctx // Create enhanced recording session to piggy-back on. tt.cgroupID, err = tt.enhancedRecorder.OpenSession(tt.ctx) require.NoError(t, err) require.Equal(t, tt.cgroupID > 0, true) - deny := []api.AddressCondition{} - allow := []api.AddressCondition{} + var deny []api.AddressCondition + var allow []api.AddressCondition for _, r := range testRanges { if len(r.deny) > 0 { deny = append(deny, api.AddressCondition{CIDR: r.deny}) @@ -294,26 +301,27 @@ func (tt *bpfContext) sendExpectDeny(t *testing.T, ver int, ip string) { } func (tt *bpfContext) expectedAuditEvent(ver int, ip string, op apievents.SessionNetwork_NetworkOperation) apievents.AuditEvent { + sessionID := tt.ctx.SessionID() return &apievents.SessionNetwork{ Metadata: apievents.Metadata{ Type: events.SessionNetworkEvent, Code: events.SessionNetworkCode, }, ServerMetadata: apievents.ServerMetadata{ - ServerID: tt.ctx.ServerID, - ServerNamespace: tt.ctx.Namespace, + ServerID: tt.ctx.GetServer().HostUUID(), + ServerNamespace: tt.ctx.GetServer().GetNamespace(), }, SessionMetadata: apievents.SessionMetadata{ - SessionID: tt.ctx.SessionID, + SessionID: sessionID.String(), }, UserMetadata: apievents.UserMetadata{ - User: tt.ctx.User, - Login: tt.ctx.Login, + User: tt.ctx.Identity.TeleportUser, + Login: tt.ctx.Identity.Login, }, BPFMetadata: apievents.BPFMetadata{ CgroupID: tt.cgroupID, Program: "restrictedsessi", - PID: uint64(tt.ctx.PID), + PID: uint64(tt.ctx.GetPID()), }, DstPort: testPort, DstAddr: ip, @@ -337,7 +345,7 @@ func TestRootNetwork(t *testing.T) { expected blockAction } - tests := []testCase{} + var tests []testCase for _, r := range testRanges { for ip, expected := range r.probe { tests = append(tests, testCase{ diff --git a/lib/srv/ctx.go b/lib/srv/ctx.go index 574994d85a7..b5a823c2a1e 100644 --- a/lib/srv/ctx.go +++ b/lib/srv/ctx.go @@ -36,10 +36,8 @@ import ( apievents "github.com/gravitational/teleport/api/types/events" apiutils "github.com/gravitational/teleport/api/utils" "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/bpf" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/pam" - restricted "github.com/gravitational/teleport/lib/restrictedsession" "github.com/gravitational/teleport/lib/services" rsession "github.com/gravitational/teleport/lib/session" "github.com/gravitational/teleport/lib/srv/uacc" @@ -153,11 +151,18 @@ type Server interface { // using reverse tunnel. UseTunnel() bool - // GetBPF returns the BPF service used for enhanced session recording. - GetBPF() bpf.BPF + // OpenBPFSession will start monitoring all events within a session and + // emitting them to the Audit Log. + OpenBPFSession(ctx *ServerContext) (uint64, error) - // GetRestrictedSessionManager returns the manager for restricting user activity - GetRestrictedSessionManager() restricted.Manager + // CloseBPFSession will stop monitoring events for a particular session. + CloseBPFSession(ctx *ServerContext) error + + // OpenRestrictedSession starts enforcing restrictions for a cgroup with cgroupID + OpenRestrictedSession(ctx *ServerContext, cgroupID uint64) + + // CloseRestrictedSession stops enforcing restrictions for a cgroup with cgroupID + CloseRestrictedSession(ctx *ServerContext, cgroupID uint64) // Context returns server shutdown context Context() context.Context @@ -172,7 +177,7 @@ type Server interface { // temporary teleport users or not GetCreateHostUser() bool - // GetHostUser returns the HostUsers instance being used to manage + // GetHostUsers returns the HostUsers instance being used to manage // host user provisioning GetHostUsers() HostUsers @@ -520,6 +525,31 @@ func (c *ServerContext) GetServer() Server { return c.srv } +// StreamWriter returns the underlying stream writer for the session or an +// events.DiscardStream if the session has yet to be established. +func (c *ServerContext) StreamWriter() events.StreamWriter { + c.mu.RLock() + defer c.mu.RUnlock() + if c.session == nil { + return &events.DiscardStream{} + } + + return c.session.Recorder() +} + +// GetPID returns the PID of the Teleport process that was re-execed +// or -1 if the process has not yet completed spawning. +func (c *ServerContext) GetPID() int { + c.mu.RLock() + defer c.mu.RUnlock() + + if c.term == nil { + return -1 + } + + return c.term.PID() +} + // CreateOrJoinSession will look in the SessionRegistry for the session ID. If // no session is found, a new one is created. If one is found, it is returned. func (c *ServerContext) CreateOrJoinSession(reg *SessionRegistry) error { @@ -800,7 +830,7 @@ func (c *ServerContext) takeClosers() []io.Closer { c.mu.Lock() defer c.mu.Unlock() - closers := []io.Closer{} + var closers []io.Closer if c.term != nil { closers = append(closers, c.term) c.term = nil diff --git a/lib/srv/ctx_test.go b/lib/srv/ctx_test.go index b627eaf4ff0..25ab14b2836 100644 --- a/lib/srv/ctx_test.go +++ b/lib/srv/ctx_test.go @@ -19,14 +19,15 @@ package srv import ( "testing" + "github.com/stretchr/testify/require" + "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/services" - "github.com/stretchr/testify/require" ) func TestCheckFileCopyingAllowed(t *testing.T) { - srv := newMockServer(t) - ctx := newTestServerContext(t, srv, nil) + srv := NewMockServer(t) + ctx := NewTestServerContext(t, srv, nil) tests := []struct { name string diff --git a/lib/srv/exec_linux_test.go b/lib/srv/exec_linux_test.go index d4906685174..8507bc0da29 100644 --- a/lib/srv/exec_linux_test.go +++ b/lib/srv/exec_linux_test.go @@ -29,8 +29,9 @@ import ( "testing" "time" - "github.com/gravitational/teleport/lib/utils" "github.com/stretchr/testify/require" + + "github.com/gravitational/teleport/lib/utils" ) // TestMain will re-execute Teleport to run a command if "exec" is passed to @@ -51,7 +52,7 @@ func TestMain(m *testing.M) { } func TestOSCommandPrep(t *testing.T) { - srv := newMockServer(t) + srv := NewMockServer(t) scx := newExecServerContext(t, srv) usr, err := user.Current() @@ -134,7 +135,7 @@ func TestOSCommandPrep(t *testing.T) { // TestContinue tests if the process hangs if a continue signal is not sent // and makes sure the process continues once it has been sent. func TestContinue(t *testing.T) { - srv := newMockServer(t) + srv := NewMockServer(t) scx := newExecServerContext(t, srv) // Configure Session Context to re-exec "ls". diff --git a/lib/srv/exec_test.go b/lib/srv/exec_test.go index 1ae83a97130..81c518757dc 100644 --- a/lib/srv/exec_test.go +++ b/lib/srv/exec_test.go @@ -23,11 +23,12 @@ import ( "strconv" "testing" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" + "github.com/gravitational/teleport" apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/sshutils" - "github.com/stretchr/testify/require" - "golang.org/x/crypto/ssh" ) // TestEmitExecAuditEvent make sure the full command and exit code for a @@ -35,7 +36,7 @@ import ( func TestEmitExecAuditEvent(t *testing.T) { t.Parallel() - srv := newMockServer(t) + srv := NewMockServer(t) scx := newExecServerContext(t, srv) expectedUsr, err := user.Current() @@ -110,7 +111,7 @@ func TestLoginDefsParser(t *testing.T) { } func newExecServerContext(t *testing.T, srv Server) *ServerContext { - scx := newTestServerContext(t, srv, nil) + scx := NewTestServerContext(t, srv, nil) term, err := newLocalTerminal(scx) require.NoError(t, err) diff --git a/lib/srv/forward/sshserver.go b/lib/srv/forward/sshserver.go index a7af5ca41cf..6e5c70095f6 100644 --- a/lib/srv/forward/sshserver.go +++ b/lib/srv/forward/sshserver.go @@ -33,10 +33,8 @@ import ( "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/bpf" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/pam" - restricted "github.com/gravitational/teleport/lib/restrictedsession" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/srv" "github.com/gravitational/teleport/lib/sshutils" @@ -59,19 +57,19 @@ import ( // // To create a forwarding server and serve a single SSH connection on it: // -// serverConfig := forward.ServerConfig{ -// ... -// } -// remoteServer, err := forward.New(serverConfig) -// if err != nil { -// return nil, trace.Wrap(err) -// } -// go remoteServer.Serve() +// serverConfig := forward.ServerConfig{ +// ... +// } +// remoteServer, err := forward.New(serverConfig) +// if err != nil { +// return nil, trace.Wrap(err) +// } +// go remoteServer.Serve() // -// conn, err := remoteServer.Dial() -// if err != nil { -// return nil, trace.Wrap(err) -// } +// conn, err := remoteServer.Dial() +// if err != nil { +// return nil, trace.Wrap(err) +// } type Server struct { log *logrus.Entry @@ -423,32 +421,34 @@ func (s *Server) UseTunnel() bool { return s.useTunnel } -// GetBPF returns the BPF service used by enhanced session recording. BPF -// for the forwarding server makes no sense (it has to run on the actual -// node), so return a NOP implementation. -func (s Server) GetBPF() bpf.BPF { - return &bpf.NOP{} +// OpenBPFSession is a nop since the session must be run on the actual node +func (s *Server) OpenBPFSession(ctx *srv.ServerContext) (uint64, error) { + return 0, nil } +// CloseBPFSession is a nop since the session must be run on the actual node +func (s *Server) CloseBPFSession(ctx *srv.ServerContext) error { + return nil +} + +// OpenRestrictedSession is a nop since the session must be run on the actual node +func (s *Server) OpenRestrictedSession(ctx *srv.ServerContext, cgroupID uint64) {} + +// CloseRestrictedSession is a nop since the session must be run on the actual node +func (s *Server) CloseRestrictedSession(ctx *srv.ServerContext, cgroupID uint64) {} + // GetCreateHostUser determines whether users should be created on the // host automatically func (s *Server) GetCreateHostUser() bool { return false } -// GetHostUser returns the HostUsers instance being used to manage +// GetHostUsers returns the HostUsers instance being used to manage // host user provisioning, unimplemented for the forwarder server. func (s *Server) GetHostUsers() srv.HostUsers { return nil } -// GetRestrictedSessionManager returns a NOP manager since for a -// forwarding server it makes no sense (it has to run on the actual -// node). -func (s Server) GetRestrictedSessionManager() restricted.Manager { - return &restricted.NOP{} -} - // GetInfo returns a services.Server that represents this server. func (s *Server) GetInfo() types.Server { return &types.ServerV2{ diff --git a/lib/srv/mock.go b/lib/srv/mock.go index d04ea13dfec..4af60225879 100644 --- a/lib/srv/mock.go +++ b/lib/srv/mock.go @@ -25,6 +25,12 @@ import ( "os/user" "testing" + "github.com/gravitational/trace" + "github.com/jonboulle/clockwork" + "github.com/sirupsen/logrus" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" + "github.com/gravitational/teleport" "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" @@ -32,22 +38,15 @@ import ( "github.com/gravitational/teleport/lib/auth" "github.com/gravitational/teleport/lib/auth/testauthority" "github.com/gravitational/teleport/lib/backend/lite" - "github.com/gravitational/teleport/lib/bpf" "github.com/gravitational/teleport/lib/events/eventstest" "github.com/gravitational/teleport/lib/fixtures" "github.com/gravitational/teleport/lib/pam" - restricted "github.com/gravitational/teleport/lib/restrictedsession" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/sshutils" "github.com/gravitational/teleport/lib/utils" - "github.com/gravitational/trace" - "github.com/jonboulle/clockwork" - "github.com/sirupsen/logrus" - "github.com/stretchr/testify/require" - "golang.org/x/crypto/ssh" ) -func newTestServerContext(t *testing.T, srv Server, roleSet services.RoleSet) *ServerContext { +func NewTestServerContext(t *testing.T, srv Server, roleSet services.RoleSet) *ServerContext { usr, err := user.Current() require.NoError(t, err) @@ -96,7 +95,7 @@ func newTestServerContext(t *testing.T, srv Server, roleSet services.RoleSet) *S return scx } -func newMockServer(t *testing.T) *mockServer { +func NewMockServer(t *testing.T) *MockServer { ctx := context.Background() clock := clockwork.NewFakeClock() @@ -126,14 +125,14 @@ func newMockServer(t *testing.T) *mockServer { authServer, err := auth.NewServer(authCfg, auth.WithClock(clock)) require.NoError(t, err) - return &mockServer{ + return &MockServer{ auth: authServer, MockEmitter: &eventstest.MockEmitter{}, clock: clock, } } -type mockServer struct { +type MockServer struct { *eventstest.MockEmitter auth *auth.Server component string @@ -141,54 +140,54 @@ type mockServer struct { } // ID is the unique ID of the server. -func (m *mockServer) ID() string { +func (m *MockServer) ID() string { return "testID" } // HostUUID is the UUID of the underlying host. For the forwarding // server this is the proxy the forwarding server is running in. -func (m *mockServer) HostUUID() string { +func (m *MockServer) HostUUID() string { return "testHostUUID" } // GetNamespace returns the namespace the server was created in. -func (m *mockServer) GetNamespace() string { +func (m *MockServer) GetNamespace() string { return "testNamespace" } // AdvertiseAddr is the publicly addressable address of this server. -func (m *mockServer) AdvertiseAddr() string { +func (m *MockServer) AdvertiseAddr() string { return "testAdvertiseAddr" } // Component is the type of server, forwarding or regular. -func (m *mockServer) Component() string { +func (m *MockServer) Component() string { return m.component } // PermitUserEnvironment returns if reading environment variables upon // startup is allowed. -func (m *mockServer) PermitUserEnvironment() bool { +func (m *MockServer) PermitUserEnvironment() bool { return false } // GetAccessPoint returns an AccessPoint for this cluster. -func (m *mockServer) GetAccessPoint() AccessPoint { +func (m *MockServer) GetAccessPoint() AccessPoint { return m.auth } // GetDataDir returns data directory of the server -func (m *mockServer) GetDataDir() string { +func (m *MockServer) GetDataDir() string { return "testDataDir" } // GetPAM returns PAM configuration for this server. -func (m *mockServer) GetPAM() (*pam.Config, error) { +func (m *MockServer) GetPAM() (*pam.Config, error) { return &pam.Config{}, nil } // GetClock returns a clock setup for the server -func (m *mockServer) GetClock() clockwork.Clock { +func (m *MockServer) GetClock() clockwork.Clock { if m.clock != nil { return m.clock } @@ -196,7 +195,7 @@ func (m *mockServer) GetClock() clockwork.Clock { } // GetInfo returns a services.Server that represents this server. -func (m *mockServer) GetInfo() types.Server { +func (m *MockServer) GetInfo() types.Server { hostname, err := os.Hostname() if err != nil { hostname = "localhost" @@ -220,50 +219,55 @@ func (m *mockServer) GetInfo() types.Server { } } -func (m *mockServer) TargetMetadata() apievents.ServerMetadata { +func (m *MockServer) TargetMetadata() apievents.ServerMetadata { return apievents.ServerMetadata{} } // UseTunnel used to determine if this node has connected to this cluster // using reverse tunnel. -func (m *mockServer) UseTunnel() bool { +func (m *MockServer) UseTunnel() bool { return false } -// GetBPF returns the BPF service used for enhanced session recording. -func (m *mockServer) GetBPF() bpf.BPF { - return &bpf.NOP{} - +// OpenBPFSession is a nop since the session must be run on the actual node +func (m *MockServer) OpenBPFSession(ctx *ServerContext) (uint64, error) { + return 0, nil } -// GetRestrictedSessionManager returns the manager for restricting user activity -func (m *mockServer) GetRestrictedSessionManager() restricted.Manager { - return &restricted.NOP{} +// CloseBPFSession is anop since the session must be run on the actual node +func (m *MockServer) CloseBPFSession(ctx *ServerContext) error { + return nil } +// OpenRestrictedSession is a nop since the session must be run on the actual node +func (m *MockServer) OpenRestrictedSession(ctx *ServerContext, cgroupID uint64) {} + +// CloseRestrictedSession is a nop since the session must be run on the actual node +func (m *MockServer) CloseRestrictedSession(ctx *ServerContext, cgroupID uint64) {} + // Context returns server shutdown context -func (m *mockServer) Context() context.Context { +func (m *MockServer) Context() context.Context { return context.Background() } // GetUtmpPath returns the path of the user accounting database and log. Returns empty for system defaults. -func (m *mockServer) GetUtmpPath() (utmp, wtmp string) { +func (m *MockServer) GetUtmpPath() (utmp, wtmp string) { return "test", "test" } // GetLockWatcher gets the server's lock watcher. -func (m *mockServer) GetLockWatcher() *services.LockWatcher { +func (m *MockServer) GetLockWatcher() *services.LockWatcher { return nil } // GetCreateHostUser gets whether the server allows host user creation // or not -func (m *mockServer) GetCreateHostUser() bool { +func (m *MockServer) GetCreateHostUser() bool { return false } // GetHostUsers -func (m *mockServer) GetHostUsers() HostUsers { +func (m *MockServer) GetHostUsers() HostUsers { return nil } diff --git a/lib/srv/regular/sshserver.go b/lib/srv/regular/sshserver.go index a26f1a97283..8726aafca23 100644 --- a/lib/srv/regular/sshserver.go +++ b/lib/srv/regular/sshserver.go @@ -279,14 +279,25 @@ func (s *Server) UseTunnel() bool { return s.useTunnel } -// GetBPF returns the BPF service used by enhanced session recording. -func (s *Server) GetBPF() bpf.BPF { - return s.ebpf +// OpenBPFSession will start monitoring all events within a session and +// emitting them to the Audit Log. +func (s *Server) OpenBPFSession(ctx *srv.ServerContext) (uint64, error) { + return s.ebpf.OpenSession(ctx) } -// GetRestrictedSessionManager returns the manager for restricting user activity. -func (s *Server) GetRestrictedSessionManager() restricted.Manager { - return s.restrictedMgr +// CloseBPFSession will stop monitoring events for a particular session. +func (s *Server) CloseBPFSession(ctx *srv.ServerContext) error { + return s.ebpf.CloseSession(ctx) +} + +// OpenRestrictedSession starts enforcing restrictions for a cgroup with cgroupID +func (s *Server) OpenRestrictedSession(ctx *srv.ServerContext, cgroupID uint64) { + s.restrictedMgr.OpenSession(ctx, cgroupID) +} + +// CloseRestrictedSession stops enforcing restrictions for a cgroup with cgroupID +func (s *Server) CloseRestrictedSession(ctx *srv.ServerContext, cgroupID uint64) { + s.restrictedMgr.CloseSession(ctx, cgroupID) } // GetLockWatcher gets the server's lock watcher. diff --git a/lib/srv/sess.go b/lib/srv/sess.go index 04c8eaaaf93..cceaec94350 100644 --- a/lib/srv/sess.go +++ b/lib/srv/sess.go @@ -31,7 +31,6 @@ import ( "github.com/gravitational/teleport/api/types" apievents "github.com/gravitational/teleport/api/types/events" "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/bpf" "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/events" "github.com/gravitational/teleport/lib/events/filesessions" @@ -943,32 +942,18 @@ func (s *session) startInteractive(ctx context.Context, ch ssh.Channel, scx *Ser return trace.Wrap(err) } - // Open a BPF recording session. If BPF was not configured, not available, - // or running in a recording proxy, OpenSession is a NOP. - sessionContext := &bpf.SessionContext{ - Context: scx.srv.Context(), - PID: s.term.PID(), - Emitter: s.Recorder(), - Namespace: scx.srv.GetNamespace(), - SessionID: s.id.String(), - ServerID: scx.srv.HostUUID(), - Login: scx.Identity.Login, - User: scx.Identity.TeleportUser, - Events: scx.Identity.AccessChecker.EnhancedRecordingSet(), - } - - if cgroupID, err := scx.srv.GetBPF().OpenSession(sessionContext); err != nil { + if cgroupID, err := scx.srv.OpenBPFSession(scx); err != nil { scx.Errorf("Failed to open enhanced recording (interactive) session: %v: %v.", s.id, err) return trace.Wrap(err) } else if cgroupID > 0 { // If a cgroup ID was assigned then enhanced session recording was enabled. s.setHasEnhancedRecording(true) - scx.srv.GetRestrictedSessionManager().OpenSession(sessionContext, cgroupID) + scx.srv.OpenRestrictedSession(scx, cgroupID) go func() { // Close the BPF recording session once the session is closed <-s.stopC - scx.srv.GetRestrictedSessionManager().CloseSession(sessionContext, cgroupID) - err = scx.srv.GetBPF().CloseSession(sessionContext) + scx.srv.CloseRestrictedSession(scx, cgroupID) + err = scx.srv.CloseBPFSession(scx) if err != nil { scx.Errorf("Failed to close enhanced recording (interactive) session: %v: %v.", s.id, err) } @@ -1140,18 +1125,7 @@ func (s *session) startExec(ctx context.Context, channel ssh.Channel, scx *Serve // Open a BPF recording session. If BPF was not configured, not available, // or running in a recording proxy, OpenSession is a NOP. - sessionContext := &bpf.SessionContext{ - Context: scx.srv.Context(), - PID: scx.ExecRequest.PID(), - Emitter: s.Recorder(), - Namespace: scx.srv.GetNamespace(), - SessionID: string(s.id), - ServerID: scx.srv.HostUUID(), - Login: scx.Identity.Login, - User: scx.Identity.TeleportUser, - Events: scx.Identity.AccessChecker.EnhancedRecordingSet(), - } - cgroupID, err := scx.srv.GetBPF().OpenSession(sessionContext) + cgroupID, err := scx.srv.OpenBPFSession(scx) if err != nil { scx.Errorf("Failed to open enhanced recording (exec) session: %v: %v.", scx.ExecRequest.GetCommand(), err) return trace.Wrap(err) @@ -1160,7 +1134,7 @@ func (s *session) startExec(ctx context.Context, channel ssh.Channel, scx *Serve // If a cgroup ID was assigned then enhanced session recording was enabled. if cgroupID > 0 { s.setHasEnhancedRecording(true) - scx.srv.GetRestrictedSessionManager().OpenSession(sessionContext, cgroupID) + scx.srv.OpenRestrictedSession(scx, cgroupID) } if tempUser != nil { @@ -1188,11 +1162,11 @@ func (s *session) startExec(ctx context.Context, channel ssh.Channel, scx *Serve // BPF session so everything can be recorded. time.Sleep(2 * time.Second) - scx.srv.GetRestrictedSessionManager().CloseSession(sessionContext, cgroupID) + scx.srv.CloseRestrictedSession(scx, cgroupID) // Close the BPF recording session. If BPF was not configured, not available, // or running in a recording proxy, this is simply a NOP. - err = scx.srv.GetBPF().CloseSession(sessionContext) + err = scx.srv.CloseBPFSession(scx) if err != nil { scx.Errorf("Failed to close enhanced recording (exec) session: %v: %v.", s.id, err) } diff --git a/lib/srv/sess_test.go b/lib/srv/sess_test.go index 1520d684acf..e270680a068 100644 --- a/lib/srv/sess_test.go +++ b/lib/srv/sess_test.go @@ -125,7 +125,7 @@ func TestSession_newRecorder(t *testing.T) { log: logger, registry: &SessionRegistry{ SessionRegistryConfig: SessionRegistryConfig{ - Srv: &mockServer{ + Srv: &MockServer{ component: teleport.ComponentNode, }, }, @@ -148,7 +148,7 @@ func TestSession_newRecorder(t *testing.T) { log: logger, registry: &SessionRegistry{ SessionRegistryConfig: SessionRegistryConfig{ - Srv: &mockServer{ + Srv: &MockServer{ component: teleport.ComponentNode, }, }, @@ -171,7 +171,7 @@ func TestSession_newRecorder(t *testing.T) { log: logger, registry: &SessionRegistry{ SessionRegistryConfig: SessionRegistryConfig{ - Srv: &mockServer{ + Srv: &MockServer{ component: teleport.ComponentNode, }, }, @@ -179,7 +179,7 @@ func TestSession_newRecorder(t *testing.T) { }, sctx: &ServerContext{ SessionRecordingConfig: nodeRecording, - srv: &mockServer{ + srv: &MockServer{ component: teleport.ComponentNode, }, }, @@ -193,7 +193,7 @@ func TestSession_newRecorder(t *testing.T) { log: logger, registry: &SessionRegistry{ SessionRegistryConfig: SessionRegistryConfig{ - Srv: &mockServer{ + Srv: &MockServer{ component: teleport.ComponentNode, }, }, @@ -201,7 +201,7 @@ func TestSession_newRecorder(t *testing.T) { }, sctx: &ServerContext{ SessionRecordingConfig: nodeRecordingSync, - srv: &mockServer{ + srv: &MockServer{ component: teleport.ComponentNode, }, Identity: IdentityContext{ @@ -231,7 +231,7 @@ func TestSession_newRecorder(t *testing.T) { log: logger, registry: &SessionRegistry{ SessionRegistryConfig: SessionRegistryConfig{ - Srv: &mockServer{ + Srv: &MockServer{ component: teleport.ComponentNode, }, }, @@ -240,7 +240,7 @@ func TestSession_newRecorder(t *testing.T) { sctx: &ServerContext{ ClusterName: "test", SessionRecordingConfig: nodeRecordingSync, - srv: &mockServer{ + srv: &MockServer{ component: teleport.ComponentNode, }, Identity: IdentityContext{ @@ -275,7 +275,7 @@ func TestSession_newRecorder(t *testing.T) { log: logger, registry: &SessionRegistry{ SessionRegistryConfig: SessionRegistryConfig{ - Srv: &mockServer{ + Srv: &MockServer{ component: teleport.ComponentNode, }, }, @@ -284,7 +284,7 @@ func TestSession_newRecorder(t *testing.T) { sctx: &ServerContext{ ClusterName: "test", SessionRecordingConfig: nodeRecordingSync, - srv: &mockServer{ + srv: &MockServer{ MockEmitter: &eventstest.MockEmitter{}, }, }, @@ -315,7 +315,7 @@ func TestSession_emitAuditEvent(t *testing.T) { }) t.Run("FallbackConcurrency", func(t *testing.T) { - srv := newMockServer(t) + srv := NewMockServer(t) reg, err := NewSessionRegistry(SessionRegistryConfig{ Srv: srv, SessionTrackerService: srv.auth, @@ -328,7 +328,7 @@ func TestSession_emitAuditEvent(t *testing.T) { log: logger, recorder: &mockRecorder{done: true}, registry: reg, - scx: newTestServerContext(t, srv, nil), + scx: NewTestServerContext(t, srv, nil), } controlCh := make(chan struct{}) @@ -355,7 +355,7 @@ func TestSession_emitAuditEvent(t *testing.T) { // Multiple sessions are opened in parallel tests to test for // deadlocks between session registry, sessions, and parties. func TestInteractiveSession(t *testing.T) { - srv := newMockServer(t) + srv := NewMockServer(t) srv.component = teleport.ComponentNode reg, err := NewSessionRegistry(SessionRegistryConfig{ @@ -384,7 +384,7 @@ func TestInteractiveSession(t *testing.T) { // TestStopUnstarted tests that a session may be stopped before it launches. func TestStopUnstarted(t *testing.T) { modules.SetTestModules(t, &modules.TestModules{TestBuildType: modules.BuildEnterprise, TestFeatures: modules.Features{ModeratedSessions: true}}) - srv := newMockServer(t) + srv := NewMockServer(t) srv.component = teleport.ComponentNode reg, err := NewSessionRegistry(SessionRegistryConfig{ @@ -426,7 +426,7 @@ func TestStopUnstarted(t *testing.T) { func TestParties(t *testing.T) { t.Parallel() - srv := newMockServer(t) + srv := NewMockServer(t) srv.component = teleport.ComponentNode // Use a separate clock from srv so we can use BlockUntil. @@ -498,7 +498,7 @@ func TestParties(t *testing.T) { } func testJoinSession(t *testing.T, reg *SessionRegistry, sess *session) { - scx := newTestServerContext(t, reg.Srv, nil) + scx := NewTestServerContext(t, reg.Srv, nil) scx.setSession(sess) // Open a new session @@ -532,7 +532,7 @@ func TestSessionRecordingModes(t *testing.T) { }, } { t.Run(tt.desc, func(t *testing.T) { - srv := newMockServer(t) + srv := NewMockServer(t) srv.component = teleport.ComponentNode reg, err := NewSessionRegistry(SessionRegistryConfig{ @@ -601,7 +601,7 @@ func TestSessionRecordingModes(t *testing.T) { } func testOpenSession(t *testing.T, reg *SessionRegistry, roleSet services.RoleSet) (*session, ssh.Channel) { - scx := newTestServerContext(t, reg.Srv, roleSet) + scx := NewTestServerContext(t, reg.Srv, roleSet) // Open a new session sshChanOpen := newMockSSHChannel()