Refactor bpf/restrictedsession usage in lib/srv (#15824)

Invert the relationship of `lib/srv` and `lib/bpf`,
`lib/restrictedsession` such that `lib/bpf` is only imported in
`lib/srv/regular`. Since not everything is built with the `bpf`
tag it's important to reduce the surface area of `lib/bpf` such
that it isn't inadvertantly imported. For instance it was entirely
possible to import a package in `tsh` that transitively
depends on `lib/bpf` - which breaks the build since `tsh` is not
compiled with the `bpf` tag.

By refactoring the `srv.Server` interface not to use the `bpf.BPF`
and `restrictedsession.Manager` interfaces directly anything that
imports `lib/srv` now won't require that `-tags=bpf` is set in
order to compile.
This commit is contained in:
rosstimothy
2022-09-07 14:49:28 +00:00
committed by GitHub
parent f54a8263f3
commit 769951fe4b
17 changed files with 303 additions and 289 deletions
+46 -39
View File
@@ -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.")
}
}
+26 -16
View File
@@ -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.
+6 -44
View File
@@ -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
}
+1 -1
View File
@@ -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 {
+13 -9
View File
@@ -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
+6 -6
View File
@@ -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 {
+8 -7
View File
@@ -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.")
}
}
+36 -28
View File
@@ -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{
+38 -8
View File
@@ -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
+4 -3
View File
@@ -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
+4 -3
View File
@@ -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".
+5 -4
View File
@@ -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)
+27 -27
View File
@@ -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{
+40 -36
View File
@@ -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
}
+17 -6
View File
@@ -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.
+8 -34
View File
@@ -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)
}
+18 -18
View File
@@ -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()