mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
+46
-39
@@ -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
@@ -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
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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.")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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".
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user