diff --git a/activator/activator.go b/activator/activator.go index 4550e2a..aa59934 100644 --- a/activator/activator.go +++ b/activator/activator.go @@ -47,10 +47,11 @@ type Server struct { forwardToTarget bool targetAddr string kubeletAddr *netip.Addr + lns Listeners } type ConnHook func(net.Conn) (conn net.Conn, cont bool, err error) -type RestoreHook func() error +type RestoreHook func() (int, error) type Option func(s *Server) @@ -60,13 +61,15 @@ func SetTargetAddr(addr string) Option { } } -func NewServer(ctx context.Context, nn ns.NetNS, opts ...Option) (*Server, error) { +func NewServer(ctx context.Context, nn ns.NetNS, connHook ConnHook, restoreHook RestoreHook, opts ...Option) (*Server, error) { s := &Server{ quit: make(chan any), connectTimeout: time.Second * 5, proxyTimeout: time.Second * 5, ns: nn, sandboxPid: parsePidFromNetNS(nn), + connHook: connHook, + restoreHook: restoreHook, } return s, nil } @@ -95,10 +98,19 @@ var ( DefaultIfaces = []string{IfaceLoopback, IfaceETH0} ) -func (s *Server) Start(ctx context.Context, connHook ConnHook, restoreHook RestoreHook, ports ...uint16) error { - s.connHook = connHook - s.restoreHook = restoreHook - s.ports = ports +func (s *Server) Start(ctx context.Context, pid int, listeners Listeners, skipStart bool) error { + s.ports = listeners.Ports() + s.lns = listeners + if !skipStart { + // populate listeners for storing them in the checkpoint. This is unused in + // this activator but is useful for forwards-compatibility when switching to + // the reuse activator. + lns, err := GetListenersOfPID(ctx, pid) + if err != nil { + return err + } + s.lns = lns + } if err := s.loadPinnedMaps(); err != nil { return err @@ -122,10 +134,20 @@ func (s *Server) Start(ctx context.Context, connHook ConnHook, restoreHook Resto } } + if skipStart { + if err := s.Reset(); err != nil { + return err + } + } + s.started = true return nil } +func (s *Server) GetListeners() []Listener { + return s.lns +} + const AttachActivatorFlag = "-zeropod-attach-activator" // AttachExec attaches the activator using exec on itself. @@ -176,12 +198,13 @@ func (s *Server) SetPeekBufferSize(size int) { // ForwardToTarget instructs the activator to forward any incoming traffic to // the specified address. The connHook and restoreHook will both be disabled. -func (s *Server) ForwardToTarget(addr string) { +func (s *Server) ForwardToTarget(_ context.Context, addr string) error { // disable hooks s.connHook = func(c net.Conn) (net.Conn, bool, error) { return c, true, nil } - s.restoreHook = func() error { return nil } + s.restoreHook = func() (int, error) { return 0, nil } s.targetAddr = addr s.forwardToTarget = true + return nil } func (s *Server) listen(ctx context.Context, port uint16) (int, error) { @@ -295,7 +318,7 @@ func (s *Server) handleConnection(ctx context.Context, netConn net.Conn, port ui } }() - if err := s.restoreHook(); err != nil { + if _, err := s.restoreHook(); err != nil { log.G(ctx).Errorf("restoreHook: %s", err) return } @@ -510,7 +533,7 @@ func (s *Server) LastActivity(port uint16) (time.Time, error) { return time.Time{}, NoActivityRecordedErr{} } - return convertBPFTime(val) + return ConvertBPFTime(val) } func (s *Server) initActivityTracker() error { @@ -527,9 +550,9 @@ func netNSPath(pid int) string { return fmt.Sprintf("/proc/%d/ns/net", pid) } -// convertBPFTime takes the value of bpf_ktime_get_ns and converts it to a +// ConvertBPFTime takes the value of bpf_ktime_get_ns and converts it to a // time.Time. -func convertBPFTime(t uint64) (time.Time, error) { +func ConvertBPFTime(t uint64) (time.Time, error) { b, err := getBootTimeNS() if err != nil { return time.Time{}, err diff --git a/activator/activator_test.go b/activator/activator_test.go index 0a42b23..4ea40be 100644 --- a/activator/activator_test.go +++ b/activator/activator_test.go @@ -146,7 +146,12 @@ func TestActivator(t *testing.T) { for name, tc := range tests { t.Run(name, func(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) - s, err := NewServer(ctx, nn) + if tc.connHook == nil { + tc.connHook = func(c net.Conn) (net.Conn, bool, error) { + return c, true, nil + } + } + s, err := NewServer(ctx, nn, tc.connHook, func() (int, error) { return 0, nil }) require.NoError(t, err) port, err := freePort() @@ -233,56 +238,52 @@ func startServer(t *testing.T, ctx context.Context, s *Server, port uint16, tc * ts := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, response) })) - if tc.connHook == nil { - tc.connHook = func(c net.Conn) (net.Conn, bool, error) { - return c, true, nil - } - } once := sync.Once{} loopIterations := 0 - err := s.Start( - ctx, - tc.connHook, - func() error { - if tc.loopConnection { - loopIterations += 1 - if loopIterations > 10 { - t.Error("loop detection failed") - return fmt.Errorf("loop detection failed") - } - // return nil + s.restoreHook = func() (int, error) { + if tc.loopConnection { + loopIterations += 1 + if loopIterations > 10 { + t.Error("loop detection failed") + return 0, fmt.Errorf("loop detection failed") } - once.Do(func() { - // simulate a delay until our server is started - time.Sleep(time.Millisecond * 200) - network := "tcp4" - if tc.ipv6 { - network = "tcp6" - } - l, err := net.Listen(network, fmt.Sprintf(":%d", port)) - require.NoError(t, err) + // return nil + } + once.Do(func() { + // simulate a delay until our server is started + time.Sleep(time.Millisecond * 200) + network := "tcp4" + if tc.ipv6 { + network = "tcp6" + } + l, err := net.Listen(network, fmt.Sprintf(":%d", port)) + require.NoError(t, err) - if !tc.loopConnection { - if err := s.DisableRedirects(); err != nil { - t.Errorf("could not disable redirects: %s", err) - } + if !tc.loopConnection { + if err := s.DisableRedirects(); err != nil { + t.Errorf("could not disable redirects: %s", err) } + } - // replace listener of server - ts.Listener.Close() - ts.Listener = l - ts.Start() - t.Logf("listening on %s", l.Addr().String()) + // replace listener of server + ts.Listener.Close() + ts.Listener = l + ts.Start() + t.Logf("listening on %s", l.Addr().String()) - t.Cleanup(func() { - l.Close() - ts.Close() - }) + t.Cleanup(func() { + l.Close() + ts.Close() }) - return nil - }, - port, + }) + return 0, nil + } + err := s.Start( + ctx, + os.Getpid(), + Listeners{{Port: port}}, + false, ) require.NoError(t, err) s.enableRedirect(port) diff --git a/activator/interface.go b/activator/interface.go new file mode 100644 index 0000000..b044043 --- /dev/null +++ b/activator/interface.go @@ -0,0 +1,20 @@ +package activator + +import ( + "context" + "time" +) + +type Activator interface { + Start(ctx context.Context, pid int, listeners Listeners, skipStart bool) error + Started() bool + Reset() error + DisableRedirects() error + AttachExec() error + SetProxyTimeout(d time.Duration) + SetConnectTimeout(d time.Duration) + LastActivity(port uint16) (time.Time, error) + Stop(ctx context.Context) + GetListeners() []Listener + ForwardToTarget(ctx context.Context, addr string) error +} diff --git a/activator/net.go b/activator/net.go new file mode 100644 index 0000000..8fd0a11 --- /dev/null +++ b/activator/net.go @@ -0,0 +1,268 @@ +package activator + +import ( + "context" + "errors" + "fmt" + "maps" + "os" + "path/filepath" + "slices" + "strconv" + "strings" + + "github.com/containerd/log" + "github.com/prometheus/procfs" + "golang.org/x/sys/unix" +) + +type Network string + +const ( + NetworkTCPAny Network = "tcp" + NetworkTCP4 Network = "tcp4" + NetworkTCP6ONLY Network = "tcp6" +) + +type Listener struct { + Port uint16 `json:"port"` + Network Network `json:"network"` + UID uint64 `json:"uid"` + Inode uint32 `json:"-"` + OrigFd int `json:"-"` + FD *os.File `json:"-"` + ownsFD bool +} + +type Listeners []Listener + +func (lns Listeners) Ports() []uint16 { + ports := map[uint16]struct{}{} + for _, ln := range lns { + ports[ln.Port] = struct{}{} + } + return slices.Collect(maps.Keys(ports)) +} + +func (ln Listener) OwnsFD() bool { + return ln.ownsFD +} + +var ErrNoListeningSockets = errors.New("no listening sockets found") + +// GetListenersOfPID gets all [Listeners] in the pid namespace. +func GetListenersOfPID(ctx context.Context, pid int, ignoredInodes ...uint64) (Listeners, error) { + return getListenersOfPID(ctx, pid, true, ignoredInodes...) +} + +// GetListenersOfPIDWithFD gets all [Listeners] in the pid namespace. +// It's the callers responsibility to close the returned listener FDs. +func GetListenersOfPIDWithFD(ctx context.Context, pid int, ignoredInodes ...uint64) (Listeners, error) { + return getListenersOfPID(ctx, pid, false, ignoredInodes...) +} + +func getListenersOfPID(ctx context.Context, pid int, closeFD bool, ignoredInodes ...uint64) (Listeners, error) { + fs, err := procfs.NewFS("/proc/" + strconv.Itoa(pid)) + if err != nil { + return nil, err + } + + netTCP4, err := fs.NetTCP() + if err != nil { + return nil, err + } + netTCP6, err := fs.NetTCP6() + if err != nil { + return nil, err + } + + listeners := Listeners{} + const tcpListen = 10 + for _, sock := range netTCP4 { + if sock.St == tcpListen { + if slices.Contains(ignoredInodes, uint64(sock.Inode)) { + continue + } + listeners = append(listeners, Listener{ + Port: uint16(sock.LocalPort), + Network: NetworkTCP4, + Inode: uint32(sock.Inode), + UID: sock.UID, + }) + } + } + for _, sock := range netTCP6 { + if sock.St == tcpListen { + if slices.Contains(ignoredInodes, uint64(sock.Inode)) { + continue + } + listeners = append(listeners, Listener{ + Port: uint16(sock.LocalPort), + Network: NetworkTCPAny, + Inode: uint32(sock.Inode), + UID: sock.UID, + }) + } + } + + pids, err := ContainerPids(pid) + if err != nil { + return nil, err + } + + inos := map[uint32]struct{}{} + for _, cpid := range pids { + for i, listener := range listeners { + if _, ok := inos[listener.Inode]; ok { + continue + } + if slices.Contains(ignoredInodes, uint64(listener.Inode)) { + continue + } + target, err := socketFdNum(cpid, []uint32{listener.Inode}) + if err != nil { + log.G(ctx).WithError(err).Debug("getting socket fd") + continue + } + pidfd, err := unix.PidfdOpen(cpid, 0) + if err != nil { + log.G(ctx).WithError(err).Debug("pidfdopen") + continue + } + defer unix.Close(pidfd) + + fd, err := unix.PidfdGetfd(pidfd, target, 0) + if err != nil { + log.G(ctx).WithError(err).Debug("pidfdgetfd") + continue + } + + sockaddr, err := unix.Getsockname(fd) + if err != nil { + log.G(ctx).WithError(err).Debug("getsockname") + _ = unix.Close(fd) + continue + } + var port int + switch sa := sockaddr.(type) { + case *unix.SockaddrInet4: + port = sa.Port + case *unix.SockaddrInet6: + port = sa.Port + } + if listener.Port != uint16(port) { + _ = unix.Close(fd) + continue + } + // the network we get from the initial fs.NetTCP is inaccurate and + // does not distinguish between dual-stack and tcp6-only so we get + // it from the socket directly. + network, err := GetNetworkFromSock(fd) + if err != nil { + log.G(ctx).WithError(err).Error("getting network from sock") + _ = unix.Close(fd) + continue + } + + if closeFD { + _ = unix.Close(fd) + } else { + listeners[i].FD = os.NewFile(uintptr(fd), "") + listeners[i].ownsFD = true + } + listeners[i].Network = network + listeners[i].OrigFd = target + inos[listener.Inode] = struct{}{} + } + } + + if len(listeners) == 0 { + return nil, ErrNoListeningSockets + } + return listeners, nil +} + +// ContainerPids returns a slice of all pids in the same pidns of pid (including +// pid). +func ContainerPids(pid int) ([]int, error) { + rootProc, err := procfs.NewProc(pid) + if err != nil { + return nil, err + } + rootNs, err := rootProc.Namespaces() + if err != nil { + return nil, err + } + pidNSInode := rootNs["pid"].Inode + + pfs, err := procfs.NewDefaultFS() + if err != nil { + return nil, err + } + procs, err := pfs.AllProcs() + if err != nil { + return nil, err + } + + containerProcs := []int{} + for _, proc := range procs { + target, err := os.Readlink(filepath.Join(procfs.DefaultMountPoint, strconv.Itoa(proc.PID), "ns", "pid")) + if err != nil { + continue + } + + fields := strings.SplitN(target, ":", 2) + if len(fields) != 2 { + continue + } + + inode, err := strconv.ParseUint(strings.Trim(fields[1], "[]"), 10, 32) + if err != nil { + continue + } + + if uint32(inode) != pidNSInode { + continue + } + containerProcs = append(containerProcs, proc.PID) + } + return containerProcs, nil +} + +// GetNetworkFromSock queries the socket fd to get the [Network] of the +// listening socket. +func GetNetworkFromSock(fd int) (Network, error) { + domain, err := unix.GetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_DOMAIN) + if err != nil { + return Network(""), err + } + if domain == unix.AF_INET6 { + if v, err := unix.GetsockoptInt(fd, unix.IPPROTO_IPV6, unix.IPV6_V6ONLY); err == nil && v == 1 { + return NetworkTCP6ONLY, nil + } + return NetworkTCPAny, nil + } + return NetworkTCP4, nil +} + +// socketFdNum scans /proc//fd for the fd number backing any of the +// socket inodes. +func socketFdNum(pid int, inodes []uint32) (int, error) { + dir := fmt.Sprintf("/proc/%d/fd", pid) + ents, err := os.ReadDir(dir) + if err != nil { + return 0, err + } + want := make(map[string]bool, len(inodes)) + for _, ino := range inodes { + want[fmt.Sprintf("socket:[%d]", ino)] = true + } + for _, e := range ents { + link, err := os.Readlink(filepath.Join(dir, e.Name())) + if err != nil || !want[link] { + continue + } + return strconv.Atoi(e.Name()) + } + return 0, fmt.Errorf("no fd for socket inodes %v", inodes) +} diff --git a/activator/reuse/activator.go b/activator/reuse/activator.go new file mode 100644 index 0000000..fe7cdb3 --- /dev/null +++ b/activator/reuse/activator.go @@ -0,0 +1,476 @@ +package reuse + +import ( + "context" + "fmt" + "net" + "net/netip" + "path/filepath" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/cilium/ebpf" + "github.com/cilium/ebpf/link" + "github.com/cilium/ebpf/rlimit" + "github.com/containerd/cgroups/v3/cgroup2" + "github.com/containerd/log" + "github.com/containernetworking/plugins/pkg/ns" + "github.com/ctrox/zeropod/activator" +) + +//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc $BPF_CLANG -cflags $BPF_CFLAGS reuseport reuseport.c -- -I/headers +//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc $BPF_CLANG -cflags $BPF_CFLAGS sockopt sockopt.c +//go:generate go run github.com/cilium/ebpf/cmd/bpf2go -cc $BPF_CLANG -cflags $BPF_CFLAGS tracker tracker.c + +type Activator struct { + *Config + ports []uint16 + mu sync.Mutex + listeners map[listenerKey]*listenerGroup + wakeInodes []uint64 + log *log.Entry + ns ns.NetNS + started atomic.Bool + sockoptLink link.Link + bindV4Link link.Link + bindV6Link link.Link + sockoptObjects *sockoptObjects + trackerLink link.Link + trackerObjs *trackerObjects + cgroupsPath string + sandboxPid int + register sync.Mutex + registeredWake atomic.Bool +} + +const ( + appKey = 0 + wakeKey = 1 + probeKey = 2 +) + +func New(ctx context.Context, ns ns.NetNS, cgroupsPath string, opts ...Option) (*Activator, error) { + cfg := &Config{} + for _, opt := range opts { + opt(cfg) + } + act := &Activator{ + ns: ns, + cgroupsPath: cgroupsPath, + log: log.GetLogger(ctx), + sandboxPid: parsePidFromNetNS(ns), + listeners: make(map[listenerKey]*listenerGroup), + Config: cfg, + } + if err := act.LoadBPF(); err != nil { + return nil, fmt.Errorf("loading ebpf: %w", err) + } + act.log.Debug("activator created") + return act, nil +} + +func (act *Activator) LoadBPF() error { + if err := rlimit.RemoveMemlock(); err != nil { + return err + } + + path, err := cgroup2.PidGroupPath(act.sandboxPid) + if err != nil { + return err + } + // use filepath.Dir to get parent cgroup of the cri-container + hostCgroupPath := filepath.Join("/sys/fs/cgroup", filepath.Dir(path)) + act.log.Debugf("attaching setsockopt to cgroup %s", hostCgroupPath) + + sockoptObjs := &sockoptObjects{} + if err := loadSockoptObjects(sockoptObjs, nil); err != nil { + return fmt.Errorf("loading sockopt objects: %w", err) + } + act.sockoptObjects = sockoptObjs + sockoptLink, err := link.AttachCgroup(link.CgroupOptions{ + Path: hostCgroupPath, + Attach: ebpf.AttachCGroupSetsockopt, + Program: sockoptObjs.Setsockopt, + }) + if err != nil { + return err + } + act.sockoptLink = sockoptLink + bindV4Link, err := link.AttachCgroup(link.CgroupOptions{ + Path: hostCgroupPath, + Attach: ebpf.AttachCGroupInet4Bind, + Program: sockoptObjs.BindV4, + }) + if err != nil { + return err + } + act.bindV4Link = bindV4Link + bindV6Link, err := link.AttachCgroup(link.CgroupOptions{ + Path: hostCgroupPath, + Attach: ebpf.AttachCGroupInet6Bind, + Program: sockoptObjs.BindV6, + }) + if err != nil { + return err + } + act.bindV6Link = bindV6Link + + trackerObjs := &trackerObjects{} + if err := loadTrackerObjects(trackerObjs, nil); err != nil { + return fmt.Errorf("loading sockopt objects: %w", err) + } + act.trackerObjs = trackerObjs + + trackerLink, err := link.AttachCgroup(link.CgroupOptions{ + Path: hostCgroupPath, + Attach: ebpf.AttachCGroupInetIngress, + Program: trackerObjs.TrackIngress, + }) + if err != nil { + return err + } + act.trackerLink = trackerLink + return nil +} + +func parsePidFromNetNS(nn ns.NetNS) int { + parts := strings.Split(nn.Path(), "/") + if len(parts) < 3 { + return 0 + } + + pid, err := strconv.Atoi(parts[2]) + if err != nil { + return 0 + } + + return pid +} + +type Config struct { + trackerIgnoreLocalhost bool + probeAddr *netip.Addr + restoreHook activator.RestoreHook +} + +type Option func(cfg *Config) + +func TrackerIgnoreLocalhost(ignore bool) Option { + return func(cfg *Config) { + cfg.trackerIgnoreLocalhost = ignore + } +} + +func RestoreHook(restoreHook activator.RestoreHook) Option { + return func(cfg *Config) { + cfg.restoreHook = restoreHook + } +} + +func ProbeAddr(addr *netip.Addr) Option { + return func(cfg *Config) { + cfg.probeAddr = addr + } +} + +func (act *Activator) Start(ctx context.Context, pid int, listeners activator.Listeners, skipStart bool) error { + act.ports = listeners.Ports() + + if skipStart { + for _, ln := range listeners { + net := ln.Network + if net == "" { + net = activator.NetworkTCPAny + } + key := listenerKey{port: ln.Port, network: net} + objs := &reuseportObjects{} + if err := loadReuseportObjects(objs, &ebpf.CollectionOptions{}); err != nil { + return fmt.Errorf("loading reuseport objects: %w", err) + } + if act.probeAddr != nil { + if err := objs.ProbeAddr.Set(act.probeAddrValue()); err != nil { + return err + } + } + act.mu.Lock() + act.listeners[key] = &listenerGroup{reuse: objs} + act.mu.Unlock() + } + if err := act.Reset(); err != nil { + return err + } + } else { + before := time.Now() + if err := act.registerListeners(pid); err != nil { + act.log.WithError(err).Error("registering listeners") + return err + } + act.log.Debugf("registered listeners in %s", time.Since(before)) + } + + act.started.Store(true) + return act.initSocketTracker() +} + +func (act *Activator) Started() bool { + return act.started.Load() +} + +func (act *Activator) Stop(_ context.Context) { + act.mu.Lock() + defer act.mu.Unlock() + for _, ln := range act.listeners { + ln.wake.close() + ln.probe.close() + ln.forwarder.close() + if ln.reuse != nil { + ln.reuse.Close() + } + } + if act.sockoptObjects != nil { + act.sockoptObjects.Close() + } + if act.sockoptLink != nil { + act.sockoptLink.Close() + } + if act.bindV4Link != nil { + act.bindV4Link.Close() + } + if act.bindV6Link != nil { + act.bindV6Link.Close() + } + if act.trackerObjs != nil { + act.trackerObjs.Close() + } + if act.trackerLink != nil { + act.trackerLink.Close() + } +} + +func (act *Activator) LastActivity(port uint16) (time.Time, error) { + if !act.Started() { + return time.Time{}, nil + } + + // the old activator used uint16 so we just convert it first + puint32 := uint32(port) + var val uint64 + if err := act.trackerObjs.SocketTracker.Lookup(&puint32, &val); err != nil { + return time.Time{}, fmt.Errorf("looking up %d: %w", port, err) + } + + if val == 0 { + return time.Time{}, activator.NoActivityRecordedErr{} + } + + return activator.ConvertBPFTime(val) +} + +func (act *Activator) initSocketTracker() error { + if err := act.clearIgnoredAddrs(); err != nil { + return err + } + if err := act.clearSocketTracker(); err != nil { + return err + } + if act.trackerIgnoreLocalhost { + if err := IgnoreAddr(act.trackerObjs.IgnoredAddrs, "127.0.0.0/8"); err != nil { + return err + } + if err := IgnoreAddr(act.trackerObjs.IgnoredAddrs, "::1/128"); err != nil { + return err + } + } + if act.probeAddr != nil { + if err := IgnoreAddr(act.trackerObjs.IgnoredAddrs, act.probeAddr.String()); err != nil { + return err + } + } + for _, port := range act.ports { + val := uint64(0) + puint32 := uint32(port) + if err := act.trackerObjs.SocketTracker.Put(&puint32, &val); err != nil { + return fmt.Errorf("unable to init activity tracker for port %d: %w", port, err) + } + } + return nil +} + +func (act *Activator) Reload(opts ...Option) error { + cfg := &Config{} + for _, opt := range opts { + opt(cfg) + } + act.Config = cfg + if err := act.initSocketTracker(); err != nil { + return err + } + act.mu.Lock() + defer act.mu.Unlock() + for _, ln := range act.listeners { + if err := ln.reuse.ProbeAddr.Set(act.probeAddrValue()); err != nil { + return err + } + } + return nil +} + +func (act *Activator) clearIgnoredAddrs() error { + var key trackerIpKey + var val byte + iter := act.trackerObjs.IgnoredAddrs.Iterate() + for iter.Next(&key, &val) { + if err := act.trackerObjs.IgnoredAddrs.Delete(key); err != nil { + return err + } + } + return iter.Err() +} + +func (act *Activator) clearSocketTracker() error { + var key uint32 + var val uint64 + iter := act.trackerObjs.SocketTracker.Iterate() + for iter.Next(&key, &val) { + if err := act.trackerObjs.SocketTracker.Delete(key); err != nil { + return err + } + } + return iter.Err() +} + +func IgnoreAddr(addrMap *ebpf.Map, ip string) error { + prefix, err := netip.ParsePrefix(ip) + if err != nil { + noPrefix, err := netip.ParseAddr(ip) + if err != nil { + return err + } + prefix, err = noPrefix.Prefix(noPrefix.BitLen()) + if err != nil { + return err + } + } + key := trackerIpKey{ + Prefixlen: uint32(prefix.Bits()), + } + + addr := prefix.Addr() + if addr.Is4() { + ip4 := addr.As4() + copy(key.Addr[:4], ip4[:]) + } else { + ip6 := addr.As16() + copy(key.Addr[:], ip6[:]) + } + + var value byte = 0 + return addrMap.Put(&key, value) +} + +func (act *Activator) wake(network activator.Network) error { + closeProbe := false + if act.restoreHook != nil { + pid, err := act.restoreHook() + if err != nil { + act.log.WithError(err).Error("restore hook") + return err + } + // TODO: should retoreHook return NoCapacity? + if pid != 0 { + closeProbe = true + before := time.Now() + act.register.Lock() + if !act.registeredWake.Load() { + if err := act.registerListeners(pid); err != nil { + act.log.WithError(err).Error("registering listeners") + act.register.Unlock() + return err + } + act.registeredWake.Store(true) + act.register.Unlock() + } + act.log.Debugf("registered listeners in %s", time.Since(before)) + } + } + act.mu.Lock() + defer act.mu.Unlock() + for _, ln := range act.listeners { + ln.wake.closeListener() + if !closeProbe { + continue + } + ln.probe.closeListener() + } + // sk_reuseport/migrate only seems to migrate pending connections to the + // wake listener only when a new conn comes in. We call poke which just + // dials and immediately closes to trigger the migration. + for _, port := range act.ports { + if err := act.poke(port, network); err != nil { + act.log.WithError(err).Error("poke app listener") + } + } + return nil +} + +func (act *Activator) poke(port uint16, network activator.Network) error { + return act.ns.Do(func(nn ns.NetNS) error { + addr := fmt.Sprintf("127.0.0.1:%d", port) + if network == activator.NetworkTCPAny || network == activator.NetworkTCP6ONLY { + addr = fmt.Sprintf("[::1]:%d", port) + } + dialer := net.Dialer{Timeout: time.Second} + c, err := dialer.Dial(string(network), addr) + if err == nil { + return c.Close() + } + return err + }) +} + +func (act *Activator) Reset() error { + act.mu.Lock() + defer act.mu.Unlock() + act.registeredWake.Store(false) + for _, ln := range act.listeners { + ln.wake.closeListener() + ln.probe.closeListener() + } + act.wakeInodes = []uint64{} + for k := range act.listeners { + if err := act.ns.Do(func(nn ns.NetNS) error { + if err := act.listenWake(k.port, k.network, act.listeners[k]); err != nil { + return err + } + if err := act.listenProbe(k.port, k.network, act.listeners[k]); err != nil { + return err + } + return nil + }); err != nil { + return err + } + if err := act.attachWake(k.network, act.listeners[k]); err != nil { + return err + } + if err := act.attachProbe(act.listeners[k]); err != nil { + return err + } + } + act.log.Debugf("listening for new connections on %d listeners: %v", len(act.listeners), act.listeners) + return nil +} + +func (act *Activator) GetListeners() []activator.Listener { + listeners := []activator.Listener{} + for k, l := range act.listeners { + listeners = append(listeners, activator.Listener{ + Port: k.port, + Network: k.network, + UID: l.app.uid, + }) + } + return listeners +} diff --git a/activator/reuse/activator_test.go b/activator/reuse/activator_test.go new file mode 100644 index 0000000..ed50a11 --- /dev/null +++ b/activator/reuse/activator_test.go @@ -0,0 +1,415 @@ +package reuse + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "os" + "os/exec" + "strings" + "sync" + "syscall" + "testing" + "time" + + "github.com/cilium/ebpf" + "github.com/containerd/log" + "github.com/containernetworking/plugins/pkg/ns" + "github.com/ctrox/zeropod/activator" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "k8s.io/utils/ptr" +) + +type testCase struct { + parallelReqs int + expectedBody string + expectedCode int + cycles int + expectLastActivity bool + trackerIgnoreLocalhost bool + kubeletAddr *netip.Addr + networks []activator.Network + forwardToFunc func(t *testing.T, port int) (string, *httptest.Server) +} + +func TestReuseActivator(t *testing.T) { + if os.Getenv("IN_NET_PID_NS") == "1" { + listenAndServe(t) + return + } + + require.NoError(t, activator.MountBPFFS(activator.BPFFSPath)) + nn, err := ns.GetCurrentNS() + require.NoError(t, err) + + c := &http.Client{ + Timeout: time.Second, + Transport: &http.Transport{ + DisableKeepAlives: true, + }, + } + + tests := map[string]testCase{ + "ipv4": { + parallelReqs: 1, + cycles: 1, + expectedBody: "app", + expectedCode: http.StatusOK, + expectLastActivity: true, + networks: []activator.Network{activator.NetworkTCP4}, + }, + "ipv6 only": { + parallelReqs: 1, + cycles: 1, + expectedBody: "app", + expectedCode: http.StatusOK, + expectLastActivity: true, + networks: []activator.Network{activator.NetworkTCP6ONLY}, + }, + "100 in parallel": { + parallelReqs: 100, + cycles: 1, + expectedBody: "app", + expectedCode: http.StatusOK, + expectLastActivity: true, + networks: []activator.Network{activator.NetworkTCPAny}, + }, + "ignore activity from localhost v4": { + parallelReqs: 1, + cycles: 1, + expectedBody: "app", + expectedCode: http.StatusOK, + expectLastActivity: false, + trackerIgnoreLocalhost: true, + networks: []activator.Network{activator.NetworkTCP4}, + }, + "ignore activity from localhost v6": { + parallelReqs: 1, + cycles: 1, + expectedBody: "app", + expectedCode: http.StatusOK, + expectLastActivity: false, + trackerIgnoreLocalhost: true, + networks: []activator.Network{activator.NetworkTCPAny}, + }, + "ignore kubelet traffic ipv4": { + parallelReqs: 1, + cycles: 1, + expectedBody: "ok\n", + expectedCode: http.StatusOK, + expectLastActivity: false, + kubeletAddr: ptr.To(netip.MustParseAddr("127.0.0.1")), + networks: []activator.Network{activator.NetworkTCP4}, + }, + "ignore kubelet traffic ipv6": { + parallelReqs: 1, + cycles: 1, + expectedBody: "ok\n", + expectedCode: http.StatusOK, + expectLastActivity: false, + kubeletAddr: ptr.To(netip.MustParseAddr("::1")), + networks: []activator.Network{activator.NetworkTCPAny}, + }, + "forward": { + parallelReqs: 1, + cycles: 1, + expectedBody: "hello from another server", + expectedCode: http.StatusOK, + expectLastActivity: true, + networks: []activator.Network{activator.NetworkTCP4}, + forwardToFunc: func(t *testing.T, port int) (string, *httptest.Server) { + ts := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprint(w, "hello from another server") + })) + // usually forwarding happens to a different pod IP, to simulate + // that we just use a different localhost IP + l, err := net.Listen("tcp4", fmt.Sprintf("127.0.0.2:%d", port)) + if err != nil { + t.Fatal(err) + } + ts.Listener.Close() + ts.Listener = l + ts.Start() + return "127.0.0.2", ts + }, + }, + "cycles": { + parallelReqs: 1, + cycles: 10, + expectedBody: "app", + expectedCode: http.StatusOK, + expectLastActivity: true, + networks: []activator.Network{activator.NetworkTCPAny}, + }, + "ipv4 and ipv6": { + parallelReqs: 1, + cycles: 1, + expectedBody: "app", + expectedCode: http.StatusOK, + expectLastActivity: true, + networks: []activator.Network{activator.NetworkTCPAny}, + }, + } + wg := sync.WaitGroup{} + for name, tc := range tests { + t.Run(name, func(t *testing.T) { + port, err := freePort() + require.NoError(t, err) + require.NoError(t, log.SetLevel(log.DebugLevel.String())) + ctx, cancel := context.WithCancel(t.Context()) + + // TODO: figure out why forwardToFunc leaks fds (maybe also the forwarder itself) + if tc.forwardToFunc == nil { + defer checkFDLeaks(t)() + } + s, err := New(ctx, nn, "/sys/fs/cgroup") + require.NoError(t, err) + + cmd, err := runApp(t, port, tc.networks...) + require.NoError(t, err) + fmt.Printf("app pid %d\n", cmd.Process.Pid) + + require.NoError(t, s.Reload( + ProbeAddr(tc.kubeletAddr), + TrackerIgnoreLocalhost(tc.trackerIgnoreLocalhost), + // TODO: not sure why but the socket migration breaks when we + // run the app in the hook itself + RestoreHook(func() (int, error) { + time.Sleep(time.Millisecond * 10) + return cmd.Process.Pid, nil + }), + )) + + t.Cleanup(func() { + cancel() + }) + + listeners := activator.Listeners{} + for _, net := range tc.networks { + listeners = append(listeners, activator.Listener{Port: uint16(port), Network: net}) + } + require.NoError(t, s.Start(ctx, os.Getpid(), listeners, true)) + if tc.forwardToFunc != nil { + addr, ts := tc.forwardToFunc(t, port) + defer ts.Close() + assert.NoError(t, s.ForwardToTarget(ctx, addr)) + assert.NoError(t, s.Reload(RestoreHook(func() (int, error) { return 0, nil }))) + } + + for range tc.cycles { + for i := 0; i < tc.parallelReqs; i++ { + for _, net := range tc.networks { + wg.Go(func() { + host := "127.0.0.1" + if net == activator.NetworkTCP6ONLY || net == activator.NetworkTCPAny { + host = "[::1]" + } + + req, err := http.NewRequest(http.MethodGet, fmt.Sprintf("http://%s:%d", host, port), nil) + if !assert.NoError(t, err) { + return + } + resp, err := c.Do(req) + if !assert.NoError(t, err) { + return + } + + b, err := io.ReadAll(resp.Body) + if !assert.NoError(t, err) { + return + } + + assert.Equal(t, tc.expectedCode, resp.StatusCode) + assert.Equal(t, tc.expectedBody, string(b)) + t.Log(string(b)) + }) + } + } + wg.Wait() + assert.NoError(t, s.Reset()) + for _, ln := range s.listeners { + if err := ln.reuse.Listeners.Delete(uint32(appKey)); err != nil { + if !errors.Is(err, ebpf.ErrKeyNotExist) { + assert.NoError(t, err) + } + } + } + } + var key uint32 + var val uint64 + count := 0 + iter := s.trackerObjs.SocketTracker.Iterate() + for iter.Next(&key, &val) { + t.Logf("found %d: %d", key, val) + count++ + } + assert.Equal(t, 1, count, "one element in socket tracker map") + last, err := s.LastActivity(uint16(port)) + if tc.expectLastActivity { + assert.NoError(t, err) + assert.Less(t, time.Since(last), time.Second) + } else { + assert.Error(t, err) + assert.ErrorIs(t, err, activator.NoActivityRecordedErr{}) + } + cancel() + s.Stop(ctx) + assert.NoError(t, cmd.Process.Kill()) + _ = cmd.Wait() + }) + } +} + +func runApp(t *testing.T, port int, networks ...activator.Network) (*exec.Cmd, error) { + cmd := exec.Command(os.Args[0], "-test.run=^TestReuseActivator$") + + nets := []string{} + for _, net := range networks { + nets = append(nets, string(net)) + } + cmd.Env = append( + os.Environ(), + "IN_NET_PID_NS=1", + fmt.Sprintf("NETWORKS=%s", strings.Join(nets, ",")), + fmt.Sprintf("ADDRESS=:%d", port), + fmt.Sprintf("PORT=%d", port), + fmt.Sprintf("RESPONSE=%s", "app"), + "GODEBUG=multipathtcp=0", + ) + cmd.Stdout = os.Stdout + cmd.Stderr = os.Stderr + cmd.SysProcAttr = &syscall.SysProcAttr{ + Cloneflags: syscall.CLONE_NEWPID, + } + r, w, err := os.Pipe() + if err != nil { + t.Fatalf("failed to create pipe: %v", err) + } + defer r.Close() + cmd.ExtraFiles = []*os.File{w} + + require.NoError(t, cmd.Start()) + w.Close() + t.Cleanup(func() { + cmd.Process.Kill() + cmd.Wait() + }) + ready := make(chan struct{}) + go func() { + buf := make([]byte, 1) + r.Read(buf) + close(ready) + }() + <-ready + t.Logf("using pid %d", cmd.Process.Pid) + return cmd, nil +} + +func freePort() (int, error) { + listener, err := net.Listen("tcp", ":0") + if err != nil { + return 0, err + } + + addr, ok := listener.Addr().(*net.TCPAddr) + if !ok { + return 0, fmt.Errorf("addr is not a net.TCPAddr: %T", listener.Addr()) + } + + if err := listener.Close(); err != nil { + return 0, err + } + + return addr.Port, nil +} + +func listenAndServe(t *testing.T) { + wg := sync.WaitGroup{} + networks := strings.SplitSeq(os.Getenv("NETWORKS"), ",") + fd := uintptr(3) + for n := range networks { + fmt.Printf("listening on %s\n", n) + ln, err := net.Listen(n, os.Getenv("ADDRESS")) + if err != nil { + t.Fatalf("create listener in isolated netns: %v", err) + } + defer ln.Close() + pipe := os.NewFile(fd, "pipe") + if pipe != nil { + pipe.Write([]byte{1}) + pipe.Close() + } + fd++ + wg.Go(func() { + http.Serve(ln, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprint(w, os.Getenv("RESPONSE")) + })) + }) + } + fmt.Println("serving") + wg.Wait() + fmt.Println("wait done") +} + +func checkFDLeaks(t *testing.T) func() { + t.Helper() + before := getFDs(t) + + return func() { + t.Helper() + + after := map[string]string{} + // some fds can take a bit of time to release so we retry a bunch + for range 10 { + after = getFDs(t) + if len(after) <= len(before) { + break + } + time.Sleep(time.Millisecond * 100) + } + if len(after) > len(before) { + b, err := json.MarshalIndent(diff(before, after), "", " ") + assert.NoError(t, err) + // TODO: fail the test here eventually but it's a bit flaky + t.Logf("file descriptor leak detected! Before: %d, After: %d\nLeaked FDs: %s", + len(before), len(after), b) + } + } +} + +func getFDs(t *testing.T) map[string]string { + t.Helper() + + fdDir := "/proc/self/fd" + entries, err := os.ReadDir(fdDir) + if err != nil { + t.Fatalf("failed to read open FDs: %v", err) + } + + fds := make(map[string]string, len(entries)) + for _, entry := range entries { + target, err := os.Readlink(fdDir + "/" + entry.Name()) + if err != nil { + target = "unknown" + } + fds[entry.Name()] = target + } + return fds +} + +func diff(before, after map[string]string) map[string]string { + leaked := make(map[string]string) + for fd, target := range after { + if _, ok := before[fd]; !ok { + leaked[fd] = target + } + } + return leaked +} diff --git a/activator/reuse/forward.go b/activator/reuse/forward.go new file mode 100644 index 0000000..952217b --- /dev/null +++ b/activator/reuse/forward.go @@ -0,0 +1,190 @@ +package reuse + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "sync" + "syscall" + "time" + + "github.com/containerd/log" + "github.com/containernetworking/plugins/pkg/ns" +) + +type forwarder struct { + targetAddr string + connectTimeout time.Duration + log *log.Entry + ln net.Listener + ns ns.NetNS + quit chan struct{} +} + +// ForwardToTarget creates a TCP proxy and replaces the app listener with it to +// forward traffic to the target addr. +// TODO: the tracker does not detect this traffic for some reason. +func (act *Activator) ForwardToTarget(ctx context.Context, addr string) error { + act.log.Infof("starting forward to target %s", addr) + act.mu.Lock() + defer act.mu.Unlock() + for k, ln := range act.listeners { + fwd := &forwarder{ + targetAddr: addr, + log: act.log.WithField("component", "forwarder"), + ns: act.ns, + // TODO: parameter + connectTimeout: time.Minute, + quit: make(chan struct{}, 1), + } + if err := act.ns.Do(func(nn ns.NetNS) error { + ln, err := listenReuseport(k.port, k.network, int(ln.app.uid)) + if err != nil { + return err + } + fwd.ln = ln + return nil + }); err != nil { + return err + } + if err := act.attachNetListener(fwd.ln, appKey, ln.reuse.Listeners, ln.reuse.SelectOrMigrate, nil); err != nil { + return fmt.Errorf("registering listener: %w", err) + } + act.listeners[k].forwarder = *fwd + go fwd.serveForward(ctx, fwd.ln, k.port) + } + return nil +} + +func (fwd *forwarder) close() { + if fwd.quit != nil { + fwd.quit <- struct{}{} + } + if fwd.ln != nil { + _ = fwd.ln.Close() + } +} + +func (fwd *forwarder) serveForward(ctx context.Context, listener net.Listener, port uint16) { + wg := sync.WaitGroup{} + + for { + conn, err := listener.Accept() + if err != nil { + select { + // TODO: we need this again? Or can we just use ctx? + case <-fwd.quit: + wg.Wait() + fwd.log.Debug("quit") + return + case <-ctx.Done(): + wg.Wait() + fwd.log.Debug("context closed") + return + default: + if !errors.Is(err, net.ErrClosed) { + fwd.log.Errorf("error accepting: %s", err) + } + } + } else { + wg.Go(func() { + fwd.log.Debugf("accepting connection from %s", conn.RemoteAddr()) + fwd.handleForwardConn(ctx, conn, port) + }) + } + } +} + +func (fwd *forwarder) handleForwardConn(ctx context.Context, conn net.Conn, port uint16) { + backendConn, err := fwd.connect(ctx, port, fwd.targetAddr) + if err != nil { + log.G(ctx).Errorf("error establishing connection: %s", err) + return + } + defer backendConn.Close() + + requestContext, cancel := context.WithTimeout(ctx, time.Minute) + defer cancel() + if err := proxy(requestContext, conn, backendConn); err != nil { + log.G(ctx).Errorf("error proxying request: %s", err) + } +} + +func (fwd *forwarder) connect(ctx context.Context, port uint16, addr string) (net.Conn, error) { + var backendConn net.Conn + dialer := net.Dialer{ + Timeout: fwd.connectTimeout, + } + targetAddr, err := net.ResolveTCPAddr("tcp", addr+":0") + if err != nil { + return nil, fmt.Errorf("parsing target addr: %w", err) + } + targetAddr.Port = int(port) + // if we dial a remote target we want a smaller timeout as we might run + // into an io timeout instead of connection refused + dialer.Timeout = time.Millisecond * 10 + fwd.log.Debugf("connecting to target address %s", targetAddr.String()) + ticker := time.NewTicker(time.Millisecond) + + defer ticker.Stop() + start := time.Now() + for { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-ticker.C: + if time.Since(start) > fwd.connectTimeout { + return nil, fmt.Errorf("timeout dialing process") + } + if err := fwd.ns.Do(func(_ ns.NetNS) error { + var err error + backendConn, err = dialer.Dial("tcp", targetAddr.String()) + return err + }); err != nil { + var serr syscall.Errno + if errors.As(err, &serr) && serr == syscall.ECONNREFUSED { + // executed program might not be ready yet, so retry in a bit. + continue + } + var operr *net.OpError + if errors.As(err, &operr) && operr.Temporary() { + fwd.log.Errorf("temporary operr: %s", operr) + continue + } + return nil, fmt.Errorf("unable to connect to process: %s", err) + } + + return backendConn, nil + } + } +} + +// proxy just proxies between conn1 and conn2. +func proxy(ctx context.Context, conn1, conn2 net.Conn) error { + defer conn1.Close() + defer conn2.Close() + + errors := make(chan error, 2) + done := make(chan struct{}, 2) + go cp(done, errors, conn2, conn1) + go cp(done, errors, conn1, conn2) + + select { + case <-ctx.Done(): + return nil + case <-done: + return nil + case err := <-errors: + return err + } +} + +func cp(done chan struct{}, errors chan error, dst io.Writer, src io.Reader) { + _, err := io.Copy(dst, src) + done <- struct{}{} + if err != nil { + errors <- err + } +} diff --git a/activator/reuse/listener.go b/activator/reuse/listener.go new file mode 100644 index 0000000..2fd6d8a --- /dev/null +++ b/activator/reuse/listener.go @@ -0,0 +1,383 @@ +package reuse + +import ( + "context" + "fmt" + "net" + "os" + "syscall" + "time" + + "github.com/cilium/ebpf" + "github.com/ctrox/zeropod/activator" + "golang.org/x/sys/unix" +) + +type listenerGroup struct { + wake wakeListener + probe probeListener + app appListener + forwarder forwarder + reuse *reuseportObjects +} + +type wakeListener struct { + ln *net.TCPListener + lnFd *os.File + epollFd int + stopFd int +} + +type appListener struct { + fd int + uid uint64 +} + +type listenerKey struct { + port uint16 + network activator.Network +} + +func (wl *wakeListener) closeListener() { + var buf [8]byte + buf[0] = 1 + _, _ = unix.Write(wl.stopFd, buf[:]) + if wl.ln != nil { + _ = wl.ln.Close() + } + if wl.lnFd != nil { + _ = wl.lnFd.Close() + } + unix.Close(wl.stopFd) +} + +func (wl *wakeListener) close() { + wl.closeListener() +} + +// listenWake will call listenReuseport and store the newly created listener in +// wake. Needs to be called inside the target network namespace. +func (act *Activator) listenWake(port uint16, network activator.Network, lg *listenerGroup) error { + act.log.Debugf("listening wake: %d %s", port, network) + ln, err := listenReuseport(port, network, int(lg.app.uid)) + if err != nil { + return fmt.Errorf("wake listener: %w", err) + } + lg.wake.ln = ln + return nil +} + +func (act *Activator) attachWake(network activator.Network, lg *listenerGroup) error { + var dupFd int + var dupErr error + if err := act.attachNetListener(lg.wake.ln, wakeKey, lg.reuse.Listeners, lg.reuse.SelectOrMigrate, func(fd uintptr) { + dupFd, dupErr = syscall.Dup(int(fd)) + }); err != nil { + return err + } + if dupErr != nil { + return dupErr + } + lg.wake.lnFd = os.NewFile(uintptr(dupFd), "") + epfd, err := unix.EpollCreate1(unix.EPOLL_CLOEXEC) + if err != nil { + act.log.WithError(err).Error("epoll create") + return err + } + lg.wake.epollFd = epfd + + var stat syscall.Stat_t + if err := syscall.Fstat(int(lg.wake.lnFd.Fd()), &stat); err != nil { + return err + } + act.wakeInodes = append(act.wakeInodes, stat.Ino) + + stopFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC) + if err != nil { + return err + } + go act.watchWake(epfd, lg.wake.lnFd.Fd(), stopFd, network) + lg.wake.stopFd = stopFd + + return nil +} + +// watchWake polls the wake listener without ever accepting and calls wake as +// soon as the poll returns something. +func (act *Activator) watchWake(epfd int, fd uintptr, stopFd int, network activator.Network) { + defer func() { + _ = unix.Close(int(fd)) + _ = unix.Close(epfd) + }() + event := unix.EpollEvent{ + Events: unix.EPOLLIN | unix.EPOLLONESHOT, + Fd: int32(fd), + } + if err := unix.EpollCtl(epfd, unix.EPOLL_CTL_ADD, int(fd), &event); err != nil { + act.log.WithError(err).Error("failed to register socket with epoll") + return + } + stopEvent := unix.EpollEvent{ + Events: unix.EPOLLIN, + Fd: int32(stopFd), + } + if err := unix.EpollCtl(epfd, unix.EPOLL_CTL_ADD, stopFd, &stopEvent); err != nil { + act.log.WithError(err).Error("failed to register stopfd with epoll") + return + } + events := make([]unix.EpollEvent, 10) + for { + n, err := unix.EpollWait(epfd, events, -1) + if err != nil { + if err == unix.EINTR { + continue + } + act.log.WithError(err).Error("epoll wait failed") + break + } + + for i := range n { + evFd := int(events[i].Fd) + if evFd == stopFd { + act.log.Debug("shutdown signal received, exiting epoll") + return + } + if evFd != int(fd) { + continue + } + act.log.Info("socket activity detected, waking up") + if err := act.wake(network); err != nil { + act.log.WithError(err).Error("wake") + } + return + } + } +} + +func (act *Activator) registerListeners(pid int) error { + act.mu.Lock() + defer act.mu.Unlock() + + before := time.Now() + listeners, err := act.listenerFds(pid) + if err != nil { + return err + } + if len(listeners) < len(act.ports) { + return fmt.Errorf("%w: expected at least %d listeners, found %d", activator.ErrNoListeningSockets, len(act.ports), len(listeners)) + } + act.log.Debugf("getting listeners in %s", time.Since(before)) + + for _, l := range listeners { + if l.FD == nil { + continue + } + defer l.FD.Close() + + key := listenerKey{port: l.Port, network: l.Network} + if _, ok := act.listeners[key]; !ok { + objs := &reuseportObjects{} + if err := loadReuseportObjects(objs, &ebpf.CollectionOptions{}); err != nil { + return fmt.Errorf("loading reuseport objects: %w", err) + } + if act.probeAddr != nil { + if err := objs.ProbeAddr.Set(act.probeAddrValue()); err != nil { + return err + } + } + act.listeners[key] = &listenerGroup{ + reuse: objs, + } + } + ln := act.listeners[key] + act.log.Debugf("registering ln %d port %d net %s ino %d", l.FD.Fd(), l.Port, l.Network, l.Inode) + if err := act.attachListener(appKey, l.FD.Fd(), ln.reuse.Listeners, ln.reuse.SelectOrMigrate); err != nil { + return fmt.Errorf("registering listener: %w", err) + } + act.log.Debugf("caching port %d fd %d uid %d", l.Port, l.OrigFd, l.UID) + act.listeners[key].app = appListener{fd: l.OrigFd, uid: l.UID} + } + if len(listeners) == 0 { + return activator.ErrNoListeningSockets + } + return nil +} + +func (act *Activator) probeAddrValue() [16]byte { + if act.probeAddr == nil { + return [16]byte{} + } + var ebpfProbeAddr [16]byte + if act.probeAddr.Is4() { + a4 := act.probeAddr.As4() + copy(ebpfProbeAddr[:], a4[:]) + } else { + a16 := act.probeAddr.As16() + copy(ebpfProbeAddr[:], a16[:]) + } + return ebpfProbeAddr +} + +func (act *Activator) listenerFds(pid int) (activator.Listeners, error) { + l, err := act.listenerFdsFromCache(pid) + if err == nil && len(l) > 0 && len(l) == len(act.listeners) { + return l, nil + } + // close fds in case the cache returned partial listeners + for _, ln := range l { + if ln.FD != nil { + _ = ln.FD.Close() + } + } + listeners, err := activator.GetListenersOfPIDWithFD(act.log.Context, pid, act.wakeInodes...) + if err != nil { + return nil, err + } + + return listeners, nil +} + +// attachListener attaches select_or_migrate to the listeners reuseport group +// and puts it into the key slot. +func (act *Activator) attachListener(key uint32, lnFd uintptr, bpfMap *ebpf.Map, prog *ebpf.Program) error { + if err := unix.SetsockoptInt(int(lnFd), unix.SOL_SOCKET, + unix.SO_ATTACH_REUSEPORT_EBPF, prog.FD()); err != nil { + return fmt.Errorf("attach reuseport prog: %w", err) + } + if err := bpfMap.Update(&key, uint64(lnFd), ebpf.UpdateAny); err != nil { + return fmt.Errorf("sockarray app: %w", err) + } + return nil +} + +// attachNetListener gets the fd of the [net.Listener] and then attaches it to +// the reuseport group into the key slot. +func (act *Activator) attachNetListener(ln net.Listener, key uint32, bpfMap *ebpf.Map, prog *ebpf.Program, fdfunc func(fd uintptr)) error { + sc, err := ln.(syscall.Conn).SyscallConn() + if err != nil { + ln.Close() + return err + } + var registerErr error + if err := sc.Control(func(fd uintptr) { + registerErr = act.attachListener(key, fd, bpfMap, prog) + if registerErr == nil { + if fdfunc != nil { + fdfunc(fd) + } + } + }); err != nil { + ln.Close() + return err + } + if registerErr != nil { + ln.Close() + return registerErr + } + return nil +} + +// listenReuseport opens a TCP listener with SO_REUSEPORT +func listenReuseport(port uint16, network activator.Network, uid int) (*net.TCPListener, error) { + lc := net.ListenConfig{ + Control: func(_, _ string, c syscall.RawConn) error { + var serr error + if err := c.Control(func(fd uintptr) { + serr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1) + if serr != nil { + return + } + if serr = unix.Fchown(int(fd), uid, -1); serr != nil { + return + } + }); err != nil { + return err + } + return serr + }, + } + lc.SetMultipathTCP(false) + ln, err := lc.Listen(context.Background(), string(network), fmt.Sprintf(":%d", port)) + if err != nil { + return nil, err + } + return ln.(*net.TCPListener), nil +} + +func (act *Activator) listenerFdsFromCache(pid int) (activator.Listeners, error) { + cache := act.listeners + if len(cache) == 0 { + return nil, nil + } + + listeners := activator.Listeners{} + pids, err := activator.ContainerPids(pid) + if err != nil { + return nil, err + } + resolved := map[listenerKey]struct{}{} + for _, cpid := range pids { + // bail out early as we found all listeners + if len(resolved) == len(cache) { + break + } + + pidfd, err := unix.PidfdOpen(cpid, 0) + if err != nil { + continue + } + defer unix.Close(pidfd) + + for k, v := range cache { + if _, ok := resolved[k]; ok { + continue + } + fd, err := unix.PidfdGetfd(pidfd, v.app.fd, 0) + if err != nil { + continue + } + var stat unix.Stat_t + if err := unix.Fstat(int(fd), &stat); err != nil { + _ = unix.Close(fd) + continue + } + + sockaddr, err := unix.Getsockname(fd) + if err != nil { + _ = unix.Close(fd) + continue + } + + var port int + switch sa := sockaddr.(type) { + case *unix.SockaddrInet4: + port = sa.Port + case *unix.SockaddrInet6: + port = sa.Port + } + if k.port != uint16(port) { + _ = unix.Close(fd) + continue + } + network, err := activator.GetNetworkFromSock(fd) + if err != nil { + _ = unix.Close(fd) + continue + } + if network != k.network { + _ = unix.Close(fd) + continue + } + resolved[k] = struct{}{} + listeners = append(listeners, activator.Listener{ + Port: k.port, + Network: network, + FD: os.NewFile(uintptr(fd), ""), + OrigFd: v.app.fd, + UID: v.app.uid, + Inode: uint32(stat.Ino), + }) + } + } + return listeners, nil +} diff --git a/activator/reuse/probe.go b/activator/reuse/probe.go new file mode 100644 index 0000000..b9afce1 --- /dev/null +++ b/activator/reuse/probe.go @@ -0,0 +1,78 @@ +package reuse + +import ( + "errors" + "fmt" + "net" + "time" + + "github.com/ctrox/zeropod/activator" +) + +type probeListener struct { + ln net.Listener +} + +// listenProbe will call listenReuseport and store the newly created listener in +// probe. Needs to be called inside the target network namespace. +func (act *Activator) listenProbe(port uint16, network activator.Network, lg *listenerGroup) error { + ln, err := listenReuseport(port, network, int(lg.app.uid)) + if err != nil { + return fmt.Errorf("wake listener: %w", err) + } + lg.probe.ln = ln + return nil +} + +func (act *Activator) attachProbe(lg *listenerGroup) error { + go func() { + for { + conn, err := lg.probe.ln.Accept() + if err != nil { + if errors.Is(err, net.ErrClosed) { + break + } + act.log.WithError(err).Error("accepting probe connection") + time.Sleep(time.Millisecond * 100) + continue + } + tcpConn, ok := conn.(*net.TCPConn) + if !ok { + act.log.Errorf("probe connection is not a *net.TCPConn: %T", conn) + _ = conn.Close() + continue + } + if err := handleProbe(tcpConn); err != nil { + act.log.WithError(err).Error("handling probe") + } + } + }() + return act.attachNetListener(lg.probe.ln, probeKey, lg.reuse.Listeners, lg.reuse.SelectOrMigrate, nil) +} + +// handleProbe writes an HTTP response to conn that satisfies the kubelet and +// immediately closes the connection. It writes a raw http response to avoid +// importing net/http which inflates the shim. +func handleProbe(conn *net.TCPConn) error { + _, err := conn.Write([]byte("HTTP/1.1 200 OK\r\nServer: zeropod probe\r\nConnection: close\r\n\r\nok\n")) + if err != nil { + return fmt.Errorf("writing probe response: %w", err) + } + if err := conn.CloseWrite(); err != nil { + return fmt.Errorf("closing write probe connection: %w", err) + } + if err := conn.Close(); err != nil { + return fmt.Errorf("closing probe connection: %w", err) + } + return err +} + +func (pl *probeListener) closeListener() { + if pl.ln != nil { + _ = pl.ln.Close() + } +} + +func (pl *probeListener) close() { + pl.closeListener() +} diff --git a/activator/reuse/reuseport.c b/activator/reuse/reuseport.c new file mode 100644 index 0000000..8f552d4 --- /dev/null +++ b/activator/reuse/reuseport.c @@ -0,0 +1,86 @@ +//go:build ignore +// SPDX-License-Identifier: GPL-2.0 + +#include +#include +#include +#include +#include +#include +#include + +char __license[] SEC("license") = "Dual MIT/GPL"; + +// key 0: app listener +// key 1: wake listener +// key 2: probe listener +struct { + __uint(type, BPF_MAP_TYPE_REUSEPORT_SOCKARRAY); + __type(key, __u32); + __type(value, __u64); + __uint(max_entries, 3); +} listeners SEC(".maps"); + +#define ETH_P_IP 0x0800 +#define ETH_P_IPV6 0x86DD +#define AF_INET6 10 +#define AF_INET 2 + +volatile __u8 probe_addr[16]; + +SEC("sk_reuseport/migrate") +int select_or_migrate(struct sk_reuseport_md *md) +{ + __u32 app = 0; + __u32 wake = 1; + __u32 probe = 2; + + // if app listener is active, pass traffic directly + if (!bpf_sk_select_reuseport(md, &listeners, &app, 0)) { + return SK_PASS; + } + + // when app is down, we want to select between probe and wake traffic + // comparing the saddr to the probe_addr + bool is_probe_traffic = false; + + struct bpf_sock *migrating = md->migrating_sk; + if (migrating) { + // during migration dst_ip4/dst_ip6 is the remote address + if (migrating->family == AF_INET) { + __u32 saddr = migrating->dst_ip4; + if (saddr == *(volatile __u32 *)probe_addr) { + // bpf_printk("being migrated: %pI4", &saddr); + is_probe_traffic = true; + } + } else { + if (__builtin_memcmp(migrating->dst_ip6, (void *)probe_addr, 16) == 0) { + // bpf_printk("being migrated: %pI6", &probe_addr); + is_probe_traffic = true; + } + } + } else if (md->eth_protocol == bpf_htons(ETH_P_IP)) { + __u32 saddr = 0; + if (bpf_skb_load_bytes_relative(md, offsetof(struct iphdr, saddr), &saddr, sizeof(saddr), BPF_HDR_START_NET) == 0) { + if (saddr == *(volatile __u32 *)probe_addr) { + is_probe_traffic = true; + } + } + } else if (md->eth_protocol == bpf_htons(ETH_P_IPV6)) { + __u32 saddr6[4] = {0}; + if (bpf_skb_load_bytes_relative(md, offsetof(struct ipv6hdr, saddr), saddr6, sizeof(saddr6), BPF_HDR_START_NET) == 0) { + if (__builtin_memcmp(saddr6, (void *)probe_addr, 16) == 0) { + is_probe_traffic = true; + } + } + } + + if (is_probe_traffic && !bpf_sk_select_reuseport(md, &listeners, &probe, 0)) { + return SK_PASS; + } + + if (!bpf_sk_select_reuseport(md, &listeners, &wake, 0)) { + return SK_PASS; + } + return SK_DROP; +} diff --git a/activator/reuse/reuseport_bpfeb.go b/activator/reuse/reuseport_bpfeb.go new file mode 100644 index 0000000..4088af7 --- /dev/null +++ b/activator/reuse/reuseport_bpfeb.go @@ -0,0 +1,135 @@ +// Code generated by bpf2go; DO NOT EDIT. +//go:build mips || mips64 || ppc64 || s390x + +package reuse + +import ( + "bytes" + _ "embed" + "fmt" + "io" + + "github.com/cilium/ebpf" +) + +// loadReuseport returns the embedded CollectionSpec for reuseport. +func loadReuseport() (*ebpf.CollectionSpec, error) { + reader := bytes.NewReader(_ReuseportBytes) + spec, err := ebpf.LoadCollectionSpecFromReader(reader) + if err != nil { + return nil, fmt.Errorf("can't load reuseport: %w", err) + } + + return spec, err +} + +// loadReuseportObjects loads reuseport and converts it into a struct. +// +// The following types are suitable as obj argument: +// +// *reuseportObjects +// *reuseportPrograms +// *reuseportMaps +// +// See ebpf.CollectionSpec.LoadAndAssign documentation for details. +func loadReuseportObjects(obj interface{}, opts *ebpf.CollectionOptions) error { + spec, err := loadReuseport() + if err != nil { + return err + } + + return spec.LoadAndAssign(obj, opts) +} + +// reuseportSpecs contains maps and programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportSpecs struct { + reuseportProgramSpecs + reuseportMapSpecs + reuseportVariableSpecs +} + +// reuseportProgramSpecs contains programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportProgramSpecs struct { + SelectOrMigrate *ebpf.ProgramSpec `ebpf:"select_or_migrate"` +} + +// reuseportMapSpecs contains maps before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportMapSpecs struct { + Listeners *ebpf.MapSpec `ebpf:"listeners"` +} + +// reuseportVariableSpecs contains global variables before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportVariableSpecs struct { + ProbeAddr *ebpf.VariableSpec `ebpf:"probe_addr"` +} + +// reuseportObjects contains all objects after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportObjects struct { + reuseportPrograms + reuseportMaps + reuseportVariables +} + +func (o *reuseportObjects) Close() error { + return _ReuseportClose( + &o.reuseportPrograms, + &o.reuseportMaps, + ) +} + +// reuseportMaps contains all maps after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportMaps struct { + Listeners *ebpf.Map `ebpf:"listeners"` +} + +func (m *reuseportMaps) Close() error { + return _ReuseportClose( + m.Listeners, + ) +} + +// reuseportVariables contains all global variables after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportVariables struct { + ProbeAddr *ebpf.Variable `ebpf:"probe_addr"` +} + +// reuseportPrograms contains all programs after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportPrograms struct { + SelectOrMigrate *ebpf.Program `ebpf:"select_or_migrate"` +} + +func (p *reuseportPrograms) Close() error { + return _ReuseportClose( + p.SelectOrMigrate, + ) +} + +func _ReuseportClose(closers ...io.Closer) error { + for _, closer := range closers { + if err := closer.Close(); err != nil { + return err + } + } + return nil +} + +// Do not access this directly. +// +//go:embed reuseport_bpfeb.o +var _ReuseportBytes []byte diff --git a/activator/reuse/reuseport_bpfeb.o b/activator/reuse/reuseport_bpfeb.o new file mode 100644 index 0000000..77bae1e Binary files /dev/null and b/activator/reuse/reuseport_bpfeb.o differ diff --git a/activator/reuse/reuseport_bpfel.go b/activator/reuse/reuseport_bpfel.go new file mode 100644 index 0000000..f490f53 --- /dev/null +++ b/activator/reuse/reuseport_bpfel.go @@ -0,0 +1,135 @@ +// Code generated by bpf2go; DO NOT EDIT. +//go:build 386 || amd64 || arm || arm64 || loong64 || mips64le || mipsle || ppc64le || riscv64 || wasm + +package reuse + +import ( + "bytes" + _ "embed" + "fmt" + "io" + + "github.com/cilium/ebpf" +) + +// loadReuseport returns the embedded CollectionSpec for reuseport. +func loadReuseport() (*ebpf.CollectionSpec, error) { + reader := bytes.NewReader(_ReuseportBytes) + spec, err := ebpf.LoadCollectionSpecFromReader(reader) + if err != nil { + return nil, fmt.Errorf("can't load reuseport: %w", err) + } + + return spec, err +} + +// loadReuseportObjects loads reuseport and converts it into a struct. +// +// The following types are suitable as obj argument: +// +// *reuseportObjects +// *reuseportPrograms +// *reuseportMaps +// +// See ebpf.CollectionSpec.LoadAndAssign documentation for details. +func loadReuseportObjects(obj interface{}, opts *ebpf.CollectionOptions) error { + spec, err := loadReuseport() + if err != nil { + return err + } + + return spec.LoadAndAssign(obj, opts) +} + +// reuseportSpecs contains maps and programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportSpecs struct { + reuseportProgramSpecs + reuseportMapSpecs + reuseportVariableSpecs +} + +// reuseportProgramSpecs contains programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportProgramSpecs struct { + SelectOrMigrate *ebpf.ProgramSpec `ebpf:"select_or_migrate"` +} + +// reuseportMapSpecs contains maps before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportMapSpecs struct { + Listeners *ebpf.MapSpec `ebpf:"listeners"` +} + +// reuseportVariableSpecs contains global variables before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type reuseportVariableSpecs struct { + ProbeAddr *ebpf.VariableSpec `ebpf:"probe_addr"` +} + +// reuseportObjects contains all objects after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportObjects struct { + reuseportPrograms + reuseportMaps + reuseportVariables +} + +func (o *reuseportObjects) Close() error { + return _ReuseportClose( + &o.reuseportPrograms, + &o.reuseportMaps, + ) +} + +// reuseportMaps contains all maps after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportMaps struct { + Listeners *ebpf.Map `ebpf:"listeners"` +} + +func (m *reuseportMaps) Close() error { + return _ReuseportClose( + m.Listeners, + ) +} + +// reuseportVariables contains all global variables after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportVariables struct { + ProbeAddr *ebpf.Variable `ebpf:"probe_addr"` +} + +// reuseportPrograms contains all programs after they have been loaded into the kernel. +// +// It can be passed to loadReuseportObjects or ebpf.CollectionSpec.LoadAndAssign. +type reuseportPrograms struct { + SelectOrMigrate *ebpf.Program `ebpf:"select_or_migrate"` +} + +func (p *reuseportPrograms) Close() error { + return _ReuseportClose( + p.SelectOrMigrate, + ) +} + +func _ReuseportClose(closers ...io.Closer) error { + for _, closer := range closers { + if err := closer.Close(); err != nil { + return err + } + } + return nil +} + +// Do not access this directly. +// +//go:embed reuseport_bpfel.o +var _ReuseportBytes []byte diff --git a/activator/reuse/reuseport_bpfel.o b/activator/reuse/reuseport_bpfel.o new file mode 100644 index 0000000..dfd21b6 Binary files /dev/null and b/activator/reuse/reuseport_bpfel.o differ diff --git a/activator/reuse/sockopt.c b/activator/reuse/sockopt.c new file mode 100644 index 0000000..ea821e7 --- /dev/null +++ b/activator/reuse/sockopt.c @@ -0,0 +1,58 @@ +//go:build ignore + +#include +#include +#include +#include + +char __license[] SEC("license") = "Dual MIT/GPL"; + +#define SOL_SOCKET 1 +#define SO_REUSEPORT 15 +#define SOL_IPV6 41 +#define IPV6_V6ONLY 26 + +SEC("cgroup/setsockopt") +int setsockopt(struct bpf_sockopt *ctx) +{ + int *optval = ctx->optval; + struct bpf_sock *sk = ctx->sk; + + if (!optval || (void *)(optval + 1) > ctx->optval_end) { + return 1; + } + + if (!sk) + return 1; + + if (sk->protocol != IPPROTO_TCP) + return 1; + + if (ctx->level == SOL_SOCKET && ctx->optname == SO_REUSEPORT) { + if (*optval == 0) { + // bpf_printk("enabling SO_REUSEPORT"); + *optval = 1; + } + } + + return 1; +} + +static __always_inline int force_so_reuseport(struct bpf_sock_addr *ctx) +{ + int reuseport_value = 1; + bpf_setsockopt(ctx, SOL_SOCKET, SO_REUSEPORT, &reuseport_value, sizeof(reuseport_value)); + return 1; +} + +SEC("cgroup/bind4") +int bind_v4(struct bpf_sock_addr *ctx) +{ + return force_so_reuseport(ctx); +} + +SEC("cgroup/bind6") +int bind_v6(struct bpf_sock_addr *ctx) +{ + return force_so_reuseport(ctx); +} diff --git a/activator/reuse/sockopt_bpfeb.go b/activator/reuse/sockopt_bpfeb.go new file mode 100644 index 0000000..70ef9ac --- /dev/null +++ b/activator/reuse/sockopt_bpfeb.go @@ -0,0 +1,135 @@ +// Code generated by bpf2go; DO NOT EDIT. +//go:build mips || mips64 || ppc64 || s390x + +package reuse + +import ( + "bytes" + _ "embed" + "fmt" + "io" + + "github.com/cilium/ebpf" +) + +// loadSockopt returns the embedded CollectionSpec for sockopt. +func loadSockopt() (*ebpf.CollectionSpec, error) { + reader := bytes.NewReader(_SockoptBytes) + spec, err := ebpf.LoadCollectionSpecFromReader(reader) + if err != nil { + return nil, fmt.Errorf("can't load sockopt: %w", err) + } + + return spec, err +} + +// loadSockoptObjects loads sockopt and converts it into a struct. +// +// The following types are suitable as obj argument: +// +// *sockoptObjects +// *sockoptPrograms +// *sockoptMaps +// +// See ebpf.CollectionSpec.LoadAndAssign documentation for details. +func loadSockoptObjects(obj interface{}, opts *ebpf.CollectionOptions) error { + spec, err := loadSockopt() + if err != nil { + return err + } + + return spec.LoadAndAssign(obj, opts) +} + +// sockoptSpecs contains maps and programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptSpecs struct { + sockoptProgramSpecs + sockoptMapSpecs + sockoptVariableSpecs +} + +// sockoptProgramSpecs contains programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptProgramSpecs struct { + BindV4 *ebpf.ProgramSpec `ebpf:"bind_v4"` + BindV6 *ebpf.ProgramSpec `ebpf:"bind_v6"` + Setsockopt *ebpf.ProgramSpec `ebpf:"setsockopt"` +} + +// sockoptMapSpecs contains maps before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptMapSpecs struct { +} + +// sockoptVariableSpecs contains global variables before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptVariableSpecs struct { +} + +// sockoptObjects contains all objects after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptObjects struct { + sockoptPrograms + sockoptMaps + sockoptVariables +} + +func (o *sockoptObjects) Close() error { + return _SockoptClose( + &o.sockoptPrograms, + &o.sockoptMaps, + ) +} + +// sockoptMaps contains all maps after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptMaps struct { +} + +func (m *sockoptMaps) Close() error { + return _SockoptClose() +} + +// sockoptVariables contains all global variables after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptVariables struct { +} + +// sockoptPrograms contains all programs after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptPrograms struct { + BindV4 *ebpf.Program `ebpf:"bind_v4"` + BindV6 *ebpf.Program `ebpf:"bind_v6"` + Setsockopt *ebpf.Program `ebpf:"setsockopt"` +} + +func (p *sockoptPrograms) Close() error { + return _SockoptClose( + p.BindV4, + p.BindV6, + p.Setsockopt, + ) +} + +func _SockoptClose(closers ...io.Closer) error { + for _, closer := range closers { + if err := closer.Close(); err != nil { + return err + } + } + return nil +} + +// Do not access this directly. +// +//go:embed sockopt_bpfeb.o +var _SockoptBytes []byte diff --git a/activator/reuse/sockopt_bpfeb.o b/activator/reuse/sockopt_bpfeb.o new file mode 100644 index 0000000..9d2a1ea Binary files /dev/null and b/activator/reuse/sockopt_bpfeb.o differ diff --git a/activator/reuse/sockopt_bpfel.go b/activator/reuse/sockopt_bpfel.go new file mode 100644 index 0000000..1ce0f55 --- /dev/null +++ b/activator/reuse/sockopt_bpfel.go @@ -0,0 +1,135 @@ +// Code generated by bpf2go; DO NOT EDIT. +//go:build 386 || amd64 || arm || arm64 || loong64 || mips64le || mipsle || ppc64le || riscv64 || wasm + +package reuse + +import ( + "bytes" + _ "embed" + "fmt" + "io" + + "github.com/cilium/ebpf" +) + +// loadSockopt returns the embedded CollectionSpec for sockopt. +func loadSockopt() (*ebpf.CollectionSpec, error) { + reader := bytes.NewReader(_SockoptBytes) + spec, err := ebpf.LoadCollectionSpecFromReader(reader) + if err != nil { + return nil, fmt.Errorf("can't load sockopt: %w", err) + } + + return spec, err +} + +// loadSockoptObjects loads sockopt and converts it into a struct. +// +// The following types are suitable as obj argument: +// +// *sockoptObjects +// *sockoptPrograms +// *sockoptMaps +// +// See ebpf.CollectionSpec.LoadAndAssign documentation for details. +func loadSockoptObjects(obj interface{}, opts *ebpf.CollectionOptions) error { + spec, err := loadSockopt() + if err != nil { + return err + } + + return spec.LoadAndAssign(obj, opts) +} + +// sockoptSpecs contains maps and programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptSpecs struct { + sockoptProgramSpecs + sockoptMapSpecs + sockoptVariableSpecs +} + +// sockoptProgramSpecs contains programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptProgramSpecs struct { + BindV4 *ebpf.ProgramSpec `ebpf:"bind_v4"` + BindV6 *ebpf.ProgramSpec `ebpf:"bind_v6"` + Setsockopt *ebpf.ProgramSpec `ebpf:"setsockopt"` +} + +// sockoptMapSpecs contains maps before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptMapSpecs struct { +} + +// sockoptVariableSpecs contains global variables before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type sockoptVariableSpecs struct { +} + +// sockoptObjects contains all objects after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptObjects struct { + sockoptPrograms + sockoptMaps + sockoptVariables +} + +func (o *sockoptObjects) Close() error { + return _SockoptClose( + &o.sockoptPrograms, + &o.sockoptMaps, + ) +} + +// sockoptMaps contains all maps after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptMaps struct { +} + +func (m *sockoptMaps) Close() error { + return _SockoptClose() +} + +// sockoptVariables contains all global variables after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptVariables struct { +} + +// sockoptPrograms contains all programs after they have been loaded into the kernel. +// +// It can be passed to loadSockoptObjects or ebpf.CollectionSpec.LoadAndAssign. +type sockoptPrograms struct { + BindV4 *ebpf.Program `ebpf:"bind_v4"` + BindV6 *ebpf.Program `ebpf:"bind_v6"` + Setsockopt *ebpf.Program `ebpf:"setsockopt"` +} + +func (p *sockoptPrograms) Close() error { + return _SockoptClose( + p.BindV4, + p.BindV6, + p.Setsockopt, + ) +} + +func _SockoptClose(closers ...io.Closer) error { + for _, closer := range closers { + if err := closer.Close(); err != nil { + return err + } + } + return nil +} + +// Do not access this directly. +// +//go:embed sockopt_bpfel.o +var _SockoptBytes []byte diff --git a/activator/reuse/sockopt_bpfel.o b/activator/reuse/sockopt_bpfel.o new file mode 100644 index 0000000..0061b6f Binary files /dev/null and b/activator/reuse/sockopt_bpfel.o differ diff --git a/activator/reuse/stub.go b/activator/reuse/stub.go new file mode 100644 index 0000000..c4a7035 --- /dev/null +++ b/activator/reuse/stub.go @@ -0,0 +1,15 @@ +package reuse + +import "time" + +func (act *Activator) DisableRedirects() error { + return nil +} + +func (act *Activator) AttachExec() error { + return nil +} + +func (act *Activator) SetProxyTimeout(d time.Duration) {} + +func (act *Activator) SetConnectTimeout(d time.Duration) {} diff --git a/activator/reuse/tracker.c b/activator/reuse/tracker.c new file mode 100644 index 0000000..b80990a --- /dev/null +++ b/activator/reuse/tracker.c @@ -0,0 +1,84 @@ +//go:build ignore + +#include +#include +#include +#include +#include + +char __license[] SEC("license") = "Dual MIT/GPL"; + +#define ETH_P_IP 0x0800 +#define ETH_P_IPV6 0x86DD +#define NEXTHDR_TCP 6 + +struct { + __uint(type, BPF_MAP_TYPE_HASH); + __type(key, __u32); // dest port + __type(value, __u64); // timestamp + __uint(max_entries, 128); // room for 128 ports in a container +} socket_tracker SEC(".maps"); + +struct ip_key { + __u32 prefixlen; + __u8 addr[16]; +}; + +struct { + __uint(type, BPF_MAP_TYPE_LPM_TRIE); + __type(key, struct ip_key); + __type(value, __u8); + __uint(max_entries, 16); + __uint(map_flags, BPF_F_NO_PREALLOC); +} ignored_addrs SEC(".maps"); + +SEC("cgroup_skb/ingress") +int track_ingress(struct __sk_buff *skb) { + __u16 proto = skb->protocol; + + if (proto == __bpf_constant_htons(ETH_P_IP)) { + struct iphdr ip; + if (bpf_skb_load_bytes(skb, 0, &ip, sizeof(ip)) < 0) return 1; + + if (ip.protocol == NEXTHDR_TCP) { + struct ip_key key = {}; + key.prefixlen = 32; + __builtin_memcpy(&key.addr, &ip.saddr, 4); + + if (bpf_map_lookup_elem(&ignored_addrs, &key)) { + return 1; + } + + __u16 dport; + __u32 l4_offset = ip.ihl * 4; + if (bpf_skb_load_bytes(skb, l4_offset + 2, &dport, 2) < 0) return 1; + + __u32 port_key = __bpf_ntohs(dport); + __u64 ts = bpf_ktime_get_ns(); + bpf_map_update_elem(&socket_tracker, &port_key, &ts, BPF_EXIST); + } + } + else if (proto == __bpf_constant_htons(ETH_P_IPV6)) { + struct ipv6hdr ipv6; + if (bpf_skb_load_bytes(skb, 0, &ipv6, sizeof(ipv6)) < 0) return 1; + + if (ipv6.nexthdr == NEXTHDR_TCP) { + struct ip_key key = {}; + key.prefixlen = 128; + __builtin_memcpy(&key.addr, &ipv6.saddr, 16); + + if (bpf_map_lookup_elem(&ignored_addrs, &key)) { + return 1; + } + + __u16 dport; + if (bpf_skb_load_bytes(skb, 40 + 2, &dport, 2) < 0) return 1; + + __u32 port_key = __bpf_ntohs(dport); + __u64 ts = bpf_ktime_get_ns(); + bpf_map_update_elem(&socket_tracker, &port_key, &ts, BPF_EXIST); + } + } + + return 1; +} diff --git a/activator/reuse/tracker_bpfeb.go b/activator/reuse/tracker_bpfeb.go new file mode 100644 index 0000000..a22b392 --- /dev/null +++ b/activator/reuse/tracker_bpfeb.go @@ -0,0 +1,143 @@ +// Code generated by bpf2go; DO NOT EDIT. +//go:build mips || mips64 || ppc64 || s390x + +package reuse + +import ( + "bytes" + _ "embed" + "fmt" + "io" + "structs" + + "github.com/cilium/ebpf" +) + +type trackerIpKey struct { + _ structs.HostLayout + Prefixlen uint32 + Addr [16]uint8 +} + +// loadTracker returns the embedded CollectionSpec for tracker. +func loadTracker() (*ebpf.CollectionSpec, error) { + reader := bytes.NewReader(_TrackerBytes) + spec, err := ebpf.LoadCollectionSpecFromReader(reader) + if err != nil { + return nil, fmt.Errorf("can't load tracker: %w", err) + } + + return spec, err +} + +// loadTrackerObjects loads tracker and converts it into a struct. +// +// The following types are suitable as obj argument: +// +// *trackerObjects +// *trackerPrograms +// *trackerMaps +// +// See ebpf.CollectionSpec.LoadAndAssign documentation for details. +func loadTrackerObjects(obj interface{}, opts *ebpf.CollectionOptions) error { + spec, err := loadTracker() + if err != nil { + return err + } + + return spec.LoadAndAssign(obj, opts) +} + +// trackerSpecs contains maps and programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerSpecs struct { + trackerProgramSpecs + trackerMapSpecs + trackerVariableSpecs +} + +// trackerProgramSpecs contains programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerProgramSpecs struct { + TrackIngress *ebpf.ProgramSpec `ebpf:"track_ingress"` +} + +// trackerMapSpecs contains maps before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerMapSpecs struct { + IgnoredAddrs *ebpf.MapSpec `ebpf:"ignored_addrs"` + SocketTracker *ebpf.MapSpec `ebpf:"socket_tracker"` +} + +// trackerVariableSpecs contains global variables before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerVariableSpecs struct { +} + +// trackerObjects contains all objects after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerObjects struct { + trackerPrograms + trackerMaps + trackerVariables +} + +func (o *trackerObjects) Close() error { + return _TrackerClose( + &o.trackerPrograms, + &o.trackerMaps, + ) +} + +// trackerMaps contains all maps after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerMaps struct { + IgnoredAddrs *ebpf.Map `ebpf:"ignored_addrs"` + SocketTracker *ebpf.Map `ebpf:"socket_tracker"` +} + +func (m *trackerMaps) Close() error { + return _TrackerClose( + m.IgnoredAddrs, + m.SocketTracker, + ) +} + +// trackerVariables contains all global variables after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerVariables struct { +} + +// trackerPrograms contains all programs after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerPrograms struct { + TrackIngress *ebpf.Program `ebpf:"track_ingress"` +} + +func (p *trackerPrograms) Close() error { + return _TrackerClose( + p.TrackIngress, + ) +} + +func _TrackerClose(closers ...io.Closer) error { + for _, closer := range closers { + if err := closer.Close(); err != nil { + return err + } + } + return nil +} + +// Do not access this directly. +// +//go:embed tracker_bpfeb.o +var _TrackerBytes []byte diff --git a/activator/reuse/tracker_bpfeb.o b/activator/reuse/tracker_bpfeb.o new file mode 100644 index 0000000..a0bf9ef Binary files /dev/null and b/activator/reuse/tracker_bpfeb.o differ diff --git a/activator/reuse/tracker_bpfel.go b/activator/reuse/tracker_bpfel.go new file mode 100644 index 0000000..556a682 --- /dev/null +++ b/activator/reuse/tracker_bpfel.go @@ -0,0 +1,143 @@ +// Code generated by bpf2go; DO NOT EDIT. +//go:build 386 || amd64 || arm || arm64 || loong64 || mips64le || mipsle || ppc64le || riscv64 || wasm + +package reuse + +import ( + "bytes" + _ "embed" + "fmt" + "io" + "structs" + + "github.com/cilium/ebpf" +) + +type trackerIpKey struct { + _ structs.HostLayout + Prefixlen uint32 + Addr [16]uint8 +} + +// loadTracker returns the embedded CollectionSpec for tracker. +func loadTracker() (*ebpf.CollectionSpec, error) { + reader := bytes.NewReader(_TrackerBytes) + spec, err := ebpf.LoadCollectionSpecFromReader(reader) + if err != nil { + return nil, fmt.Errorf("can't load tracker: %w", err) + } + + return spec, err +} + +// loadTrackerObjects loads tracker and converts it into a struct. +// +// The following types are suitable as obj argument: +// +// *trackerObjects +// *trackerPrograms +// *trackerMaps +// +// See ebpf.CollectionSpec.LoadAndAssign documentation for details. +func loadTrackerObjects(obj interface{}, opts *ebpf.CollectionOptions) error { + spec, err := loadTracker() + if err != nil { + return err + } + + return spec.LoadAndAssign(obj, opts) +} + +// trackerSpecs contains maps and programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerSpecs struct { + trackerProgramSpecs + trackerMapSpecs + trackerVariableSpecs +} + +// trackerProgramSpecs contains programs before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerProgramSpecs struct { + TrackIngress *ebpf.ProgramSpec `ebpf:"track_ingress"` +} + +// trackerMapSpecs contains maps before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerMapSpecs struct { + IgnoredAddrs *ebpf.MapSpec `ebpf:"ignored_addrs"` + SocketTracker *ebpf.MapSpec `ebpf:"socket_tracker"` +} + +// trackerVariableSpecs contains global variables before they are loaded into the kernel. +// +// It can be passed ebpf.CollectionSpec.Assign. +type trackerVariableSpecs struct { +} + +// trackerObjects contains all objects after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerObjects struct { + trackerPrograms + trackerMaps + trackerVariables +} + +func (o *trackerObjects) Close() error { + return _TrackerClose( + &o.trackerPrograms, + &o.trackerMaps, + ) +} + +// trackerMaps contains all maps after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerMaps struct { + IgnoredAddrs *ebpf.Map `ebpf:"ignored_addrs"` + SocketTracker *ebpf.Map `ebpf:"socket_tracker"` +} + +func (m *trackerMaps) Close() error { + return _TrackerClose( + m.IgnoredAddrs, + m.SocketTracker, + ) +} + +// trackerVariables contains all global variables after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerVariables struct { +} + +// trackerPrograms contains all programs after they have been loaded into the kernel. +// +// It can be passed to loadTrackerObjects or ebpf.CollectionSpec.LoadAndAssign. +type trackerPrograms struct { + TrackIngress *ebpf.Program `ebpf:"track_ingress"` +} + +func (p *trackerPrograms) Close() error { + return _TrackerClose( + p.TrackIngress, + ) +} + +func _TrackerClose(closers ...io.Closer) error { + for _, closer := range closers { + if err := closer.Close(); err != nil { + return err + } + } + return nil +} + +// Do not access this directly. +// +//go:embed tracker_bpfel.o +var _TrackerBytes []byte diff --git a/activator/reuse/tracker_bpfel.o b/activator/reuse/tracker_bpfel.o new file mode 100644 index 0000000..6d5260f Binary files /dev/null and b/activator/reuse/tracker_bpfel.o differ diff --git a/api/node/v1/meta.go b/api/node/v1/meta.go index 4e05669..6b94fe1 100644 --- a/api/node/v1/meta.go +++ b/api/node/v1/meta.go @@ -13,6 +13,7 @@ const ( NodeNameEnvKey = "NODE_NAME" PodIPEnvKey = "POD_IP" preDumpDirName = "pre-dump" + listenersFile = "zeropod_listeners.json" ) var imageBasePath = "/var/lib/zeropod/" @@ -48,3 +49,7 @@ func PreDumpDir(id string) string { func RelativePreDumpDir() string { return filepath.Join("..", SnapshotSuffix, preDumpDirName) } + +func ListenersFile(id string) string { + return filepath.Join(ImagePath(id), SnapshotSuffix, listenersFile) +} diff --git a/api/shim/v1/config.go b/api/shim/v1/config.go index 553ed79..7d0ff78 100644 --- a/api/shim/v1/config.go +++ b/api/shim/v1/config.go @@ -18,6 +18,7 @@ import ( ) const ( + DefaultOptDir = "/opt/zeropod" ConfigDir = "etc" ConfigFileName = "shim.json" NodeLabel = "zeropod.ctrox.dev/node" @@ -54,6 +55,7 @@ const ( DefaultProbeBinaryName = "kubelet" DefaultTrackerIgnoreLocalhost = true DefaultCapacityRequest = false + DefaultReuseportActivator = false ) var ContainerdAnnotations = []string{ @@ -95,8 +97,10 @@ type AnnotationConfig struct { } type Config struct { - TrackerIgnoreLocalhost bool `json:"trackerIgnoreLocalhost"` - CapacityRequest bool `json:"capacityRequest"` + TrackerIgnoreLocalhost bool `json:"trackerIgnoreLocalhost"` + CapacityRequest bool `json:"capacityRequest"` + ProbeAddress string `json:"probeAddress"` + ReuseportActivator bool `json:"reuseportActivator"` AnnotationConfig `json:"-"` } @@ -228,12 +232,13 @@ func NewConfig(ctx context.Context, spec *specs.Spec) (*Config, error) { cfg := &Config{ TrackerIgnoreLocalhost: DefaultTrackerIgnoreLocalhost, CapacityRequest: DefaultCapacityRequest, + ReuseportActivator: DefaultReuseportActivator, } - e, err := os.Executable() + path, err := relativeConfigFile() if err != nil { - return nil, fmt.Errorf("getting executable dir: %w", err) + return nil, err } - b, err := os.ReadFile(filepath.Join(filepath.Dir(e), "..", ConfigDir, ConfigFileName)) + b, err := os.ReadFile(path) if err == nil { if err := json.Unmarshal(b, cfg); err != nil { return nil, err @@ -264,6 +269,14 @@ func NewConfig(ctx context.Context, spec *specs.Spec) (*Config, error) { return cfg, nil } +func relativeConfigFile() (string, error) { + e, err := os.Executable() + if err != nil { + return "", fmt.Errorf("getting executable dir: %w", err) + } + return filepath.Join(filepath.Dir(e), "..", ConfigDir, ConfigFileName), nil +} + func (cfg Config) IsZeropodContainer() bool { if slices.Contains(cfg.ZeropodContainerNames, cfg.ContainerName) { return true @@ -295,3 +308,41 @@ func LiveMigrationEnabled(annotations map[string]string) bool { _, ok := annotations[LiveMigrateAnnotationKey] return ok } + +func (cfg Config) LastModified() time.Time { + configPath, err := relativeConfigFile() + if err != nil { + return time.Time{} + } + info, err := os.Stat(configPath) + if err != nil { + return time.Time{} + } + return info.ModTime() +} + +func Load(optDir string) (*Config, error) { + b, err := os.ReadFile(filepath.Join(optDir, ConfigDir, ConfigFileName)) + if err != nil { + return nil, err + } + cfg := &Config{} + if err := json.Unmarshal(b, cfg); err != nil { + return nil, err + } + return cfg, nil +} + +func (cfg *Config) Write(optPath string) error { + b, err := json.MarshalIndent(cfg, "", " ") + if err != nil { + return fmt.Errorf("marshaling config: %w", err) + } + if err := os.MkdirAll(filepath.Join(optPath, ConfigDir), os.ModePerm); err != nil { + return err + } + if err := os.WriteFile(filepath.Join(optPath, ConfigDir, ConfigFileName), b, 0600); err != nil { + return fmt.Errorf("unable to write shim file: %w", err) + } + return nil +} diff --git a/cmd/installer/main.go b/cmd/installer/main.go index 8b00862..33eebcb 100644 --- a/cmd/installer/main.go +++ b/cmd/installer/main.go @@ -4,7 +4,6 @@ import ( "bytes" "context" "crypto/x509" - "encoding/json" "encoding/pem" "errors" "flag" @@ -41,6 +40,7 @@ var ( versionFlag = flag.Bool("version", false, "output version and exit") trackerIgnoreLocalhost = flag.Bool("tracker-ignore-localhost", v1.DefaultTrackerIgnoreLocalhost, "set to ignore traffic from localhost in socket tracker") capacityRequest = flag.Bool("capacity-request", v1.DefaultCapacityRequest, "enable shim to make a capacity request before restoring") + reuseportActivator = flag.Bool("reuseport-activator", v1.DefaultReuseportActivator, "enable the new reuseport activator") //lint:ignore U1000 kept for compatibility probeBinaryName = flag.String("probe-binary-name", v1.DefaultProbeBinaryName, "Deprecated: this is no longer used, flag will be removed in future release") @@ -232,18 +232,19 @@ func installRuntime(ctx context.Context, runtime containerRuntime) error { return fmt.Errorf("unable to write shim file: %w", err) } - b, err := json.MarshalIndent(&v1.Config{ - TrackerIgnoreLocalhost: *trackerIgnoreLocalhost, - CapacityRequest: *capacityRequest, - }, "", " ") + cfg, err := v1.Load(opt) if err != nil { - return fmt.Errorf("marshaling config: %w", err) - } - if err := os.MkdirAll(filepath.Join(opt, v1.ConfigDir), os.ModePerm); err != nil { - return err - } - if err := os.WriteFile(filepath.Join(opt, v1.ConfigDir, v1.ConfigFileName), b, 0600); err != nil { - return fmt.Errorf("unable to write shim file: %w", err) + if !os.IsNotExist(err) { + return fmt.Errorf("loading config: %w", err) + } + log.Printf("existing config not found, creating from scratch") + cfg = &v1.Config{} + } + cfg.TrackerIgnoreLocalhost = *trackerIgnoreLocalhost + cfg.CapacityRequest = *capacityRequest + cfg.ReuseportActivator = *reuseportActivator + if err := cfg.Write(opt); err != nil { + return fmt.Errorf("writing config: %w", err) } if runtime == runtimeK3S { diff --git a/cmd/manager/main.go b/cmd/manager/main.go index 0d40c81..8a4aca9 100644 --- a/cmd/manager/main.go +++ b/cmd/manager/main.go @@ -155,7 +155,7 @@ func main() { EnableOpenMetrics: true, }), ) - mux.Handle("/probe", http.HandlerFunc(probeHander(redirector))) + mux.Handle("/probe", http.HandlerFunc(probeHander(redirector, log))) server := &http.Server{Addr: *metricsAddr, Handler: mux} go func() { @@ -255,14 +255,16 @@ func newControllerManager(nodeName string) (ctrlmanager.Manager, error) { // probeHandler responds to kublet probes and extracts the remoteAddr to // populate the kubeletAddr of the redirector which is used for probe detection. -func probeHander(redirector *manager.Redirector) http.HandlerFunc { +func probeHander(redirector *manager.Redirector, log *slog.Logger) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { host, _, err := net.SplitHostPort(r.RemoteAddr) if err == nil { addr, err := netip.ParseAddr(host) if err == nil { if redirector != nil { - redirector.InitKubeletAddr(addr) + if err := redirector.InitKubeletAddr(addr); err != nil { + log.Error("init kubelet addr failed", "error", err) + } } } } diff --git a/config/base/node-daemonset.yaml b/config/base/node-daemonset.yaml index 2c7d0fa..1c1b1af 100644 --- a/config/base/node-daemonset.yaml +++ b/config/base/node-daemonset.yaml @@ -83,6 +83,8 @@ spec: name: zeropod-run - mountPath: /var/lib/zeropod name: zeropod-var + - mountPath: /opt/zeropod + name: zeropod-opt - mountPath: /hostproc name: hostproc - mountPath: /sys/fs/bpf diff --git a/config/reuseport-activator/kustomization.yaml b/config/reuseport-activator/kustomization.yaml new file mode 100644 index 0000000..e1fed53 --- /dev/null +++ b/config/reuseport-activator/kustomization.yaml @@ -0,0 +1,9 @@ +apiVersion: kustomize.config.k8s.io/v1alpha1 +kind: Component +patches: + - patch: |- + - op: add + path: /spec/template/spec/initContainers/0/args/- + value: -reuseport-activator=true + target: + kind: DaemonSet diff --git a/criu/socket-uid.patch b/criu/socket-uid.patch new file mode 100644 index 0000000..8a16fc1 --- /dev/null +++ b/criu/socket-uid.patch @@ -0,0 +1,61 @@ +diff --git a/criu/include/sk-inet.h b/criu/include/sk-inet.h +index 69ee8589e..0a60fc5bb 100644 +--- a/criu/include/sk-inet.h ++++ b/criu/include/sk-inet.h +@@ -41,6 +41,7 @@ struct inet_sk_desc { + unsigned int dst_addr[4]; + unsigned short shutdown; + bool cork; ++ uid_t uid; + + int rfd; + int cpt_reuseaddr; +diff --git a/criu/sk-inet.c b/criu/sk-inet.c +index 422edc656..73277cd2d 100644 +--- a/criu/sk-inet.c ++++ b/criu/sk-inet.c +@@ -526,6 +526,8 @@ static int do_dump_one_inet_fd(int lfd, u32 id, const struct fd_parms *p, int fa + ie.dst_port = sk->dst_port; + ie.backlog = sk->wqlen; + ie.flags = p->flags; ++ ie.has_uid = true; ++ ie.uid = sk->uid; + + ie.fown = (FownEntry *)&p->fown; + ie.opts = &skopts; +@@ -678,6 +680,8 @@ int inet_collect_one(struct nlmsghdr *h, int family, int type, struct ns_id *ns) + d->wqlen = m->idiag_wqueue; + memcpy(d->src_addr, m->id.idiag_src, sizeof(u32) * 4); + memcpy(d->dst_addr, m->id.idiag_dst, sizeof(u32) * 4); ++ d->uid = m->idiag_uid; ++ pr_info("storing socket uid: %d\n", d->uid); + + if (tb[INET_DIAG_SHUTDOWN]) + d->shutdown = nla_get_u8(tb[INET_DIAG_SHUTDOWN]); +@@ -890,6 +894,16 @@ static int open_inet_sk(struct file_desc *d, int *new_fd) + return -1; + } + ++ if (ie->has_uid) { ++ pr_info("restoring socket uid: %d", ie->uid); ++ if (fchown(sk, ie->uid, -1) < 0) { ++ pr_err("Failed to set socket UID to %u", ie->uid); ++ goto err; ++ } ++ } else { ++ pr_info("socket uid not found! %d\n", ie->uid); ++ } ++ + if (reset_setsockcreatecon()) + goto err; + +diff --git a/images/sk-inet.proto b/images/sk-inet.proto +index 2c709e018..b6fc3e9d8 100644 +--- a/images/sk-inet.proto ++++ b/images/sk-inet.proto +@@ -58,4 +58,5 @@ message inet_sk_entry { + optional uint32 ns_id = 18; + optional sk_shutdown shutdown = 19; + optional tcp_opts_entry tcp_opts = 20; ++ optional uint32 uid = 21; + } diff --git a/manager/redirector_attacher.go b/manager/redirector_attacher.go index 6b1e939..e720d16 100644 --- a/manager/redirector_attacher.go +++ b/manager/redirector_attacher.go @@ -15,6 +15,7 @@ import ( "github.com/containernetworking/plugins/pkg/ns" "github.com/ctrox/zeropod/activator" + v1 "github.com/ctrox/zeropod/api/shim/v1" "github.com/fsnotify/fsnotify" ) @@ -59,12 +60,22 @@ func AttachRedirectors(ctx context.Context, log *slog.Logger, activatorOpts ...a } // InitKubeletAddr sets the kubelet addr on the [Redirector] if it's unset. -func (r *Redirector) InitKubeletAddr(addr netip.Addr) { +func (r *Redirector) InitKubeletAddr(addr netip.Addr) error { if r.kubeletAddr != nil { - return + return nil } r.log.Info("redirector kubelet addr set", "addr", addr) r.kubeletAddr = &addr + return r.storekubeletAddrInConfig() +} + +func (r *Redirector) storekubeletAddrInConfig() error { + cfg, err := v1.Load(v1.DefaultOptDir) + if err != nil { + return fmt.Errorf("lading config: %w", err) + } + cfg.ProbeAddress = r.kubeletAddr.String() + return cfg.Write(v1.DefaultOptDir) } func (r *Redirector) reconcile() error { diff --git a/shim/checkpoint.go b/shim/checkpoint.go index f2f96fd..60233bc 100644 --- a/shim/checkpoint.go +++ b/shim/checkpoint.go @@ -2,6 +2,7 @@ package shim import ( "context" + "encoding/json" "fmt" "io" "os" @@ -14,6 +15,7 @@ import ( "github.com/containerd/containerd/v2/cmd/containerd-shim-runc-v2/process" runcC "github.com/containerd/go-runc" "github.com/containerd/log" + "github.com/ctrox/zeropod/activator" nodev1 "github.com/ctrox/zeropod/api/node/v1" v1 "github.com/ctrox/zeropod/api/shim/v1" "github.com/icza/backscanner" @@ -129,6 +131,10 @@ func (c *Container) checkpoint(ctx context.Context) error { return err } + if err := c.storeListeners(); err != nil { + log.G(ctx).WithError(err).Error("storing listeners in snaphsot path") + } + c.setPhaseNotify(v1.ContainerPhase_SCALED_DOWN, time.Since(beforeCheckpoint)) log.G(ctx).Infof("checkpointing done in %s", c.metrics.LastCheckpointDuration.AsDuration()) @@ -157,6 +163,29 @@ func (c *Container) checkpointExtraArgs() []string { return def } +func (c *Container) storeListeners() error { + f, err := os.Create(nodev1.ListenersFile(c.ID())) + if err != nil { + return err + } + defer f.Close() + listeners := c.activator.GetListeners() + return json.NewEncoder(f).Encode(listeners) +} + +func (c *Container) loadListeners() (activator.Listeners, error) { + f, err := os.Open(nodev1.ListenersFile(c.ID())) + if err != nil { + return nil, err + } + defer f.Close() + listeners := activator.Listeners{} + if err := json.NewDecoder(f).Decode(&listeners); err != nil { + return nil, err + } + return listeners, nil +} + func printCriuLogs(ctx context.Context, file string) string { lines, err := getLastLines(ctx, file, 20) if err != nil { diff --git a/shim/container.go b/shim/container.go index d8d5641..8ce3350 100644 --- a/shim/container.go +++ b/shim/container.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "net/netip" "os" "slices" "sync" @@ -19,6 +20,7 @@ import ( "github.com/containerd/log" "github.com/containernetworking/plugins/pkg/ns" "github.com/ctrox/zeropod/activator" + "github.com/ctrox/zeropod/activator/reuse" nodev1 "github.com/ctrox/zeropod/api/node/v1" v1 "github.com/ctrox/zeropod/api/shim/v1" "google.golang.org/protobuf/proto" @@ -38,7 +40,7 @@ type Container struct { context context.Context id string createOpts *anypb.Any - activator *activator.Server + activator activator.Activator cfg *v1.Config initialProcess process.Process process process.Process @@ -62,6 +64,7 @@ type Container struct { evacuation sync.Once metrics *v1.ContainerMetrics runcVersion string + lastConfigReload time.Time } func New(ctx context.Context, cfg *v1.Config, r *taskAPI.CreateTaskRequest, pt stdio.Platform, events chan *v1.ContainerStatus) (*Container, error) { @@ -101,6 +104,16 @@ func New(ctx context.Context, cfg *v1.Config, r *taskAPI.CreateTaskRequest, pt s metrics: newMetrics(cfg, true), runcVersion: vers.Runc, } + + if c.cfg.ReuseportActivator { + log.G(ctx).Info("using reuseport activator") + act, err := reuse.New(ctx, c.netNS, c.cfg.Spec.Linux.CgroupsPath, c.activatorOpts(ctx)...) + if err != nil { + return nil, err + } + c.activator = act + } + return c, nil } @@ -115,7 +128,7 @@ func (c *Container) Register(ctx context.Context, container *runc.Container) err c.process = p c.initialProcess = p - if err := c.initActivator(ctx, c.SkipStart()); err != nil { + if err := c.initActivator(ctx); err != nil { log.G(ctx).Warnf("activator init failed, disabling scale down: %s", err) c.cfg.ScaleDownDuration = 0 } @@ -134,6 +147,35 @@ func (c *Container) Config() *v1.Config { return c.cfg } +func (c *Container) reloadConfig(ctx context.Context) error { + if c.cfg.LastModified().Equal(c.lastConfigReload) { + log.G(ctx).Infof("config file up to date: %s", c.lastConfigReload) + return nil + } + log.G(ctx).Infof("reloading config") + spec, err := GetSpec(c.Bundle) + if err != nil { + return fmt.Errorf("getting container spec: %w", err) + } + // copy ports since they might have been discovered on first startup + ports := c.cfg.Ports + cfg, err := v1.NewConfig(ctx, spec) + if err != nil { + return fmt.Errorf("creating config: %w", err) + } + c.cfg = cfg + if len(c.cfg.Ports) == 0 { + c.cfg.Ports = ports + } + if act, ok := c.activator.(*reuse.Activator); ok { + if err := act.Reload(c.activatorOpts(ctx)...); err != nil { + return err + } + } + c.lastConfigReload = c.cfg.LastModified() + return nil +} + func (c *Container) ScheduleScaleDown() { c.scheduleScaleDownIn(c.cfg.ScaleDownDuration) } @@ -381,21 +423,38 @@ func (c *Container) EvacDrainStarted() bool { return c.evacDrainStarted.Load() } -func (c *Container) initActivator(ctx context.Context, enableRedirects bool) error { +func (c *Container) activatorOpts(ctx context.Context) []reuse.Option { + opts := []reuse.Option{ + reuse.RestoreHook(c.restoreHandler(c.context)), + reuse.TrackerIgnoreLocalhost(c.cfg.TrackerIgnoreLocalhost), + } + if c.cfg.ProbeAddress != "" { + addr, err := netip.ParseAddr(c.cfg.ProbeAddress) + if err != nil { + log.G(ctx).WithError(err).Warn("invalid probe address configured") + } else { + opts = append(opts, reuse.ProbeAddr(&addr)) + } + } + return opts +} + +func (c *Container) initActivator(ctx context.Context) error { c.cancelInit() - if c.activator == nil { - act, err := activator.NewServer(ctx, c.netNS) + if c.activator == nil && !c.cfg.ReuseportActivator { + log.G(ctx).Info("using legacy activator") + act, err := activator.NewServer(ctx, c.netNS, c.detectProbe(ctx), c.restoreHandler(c.context)) if err != nil { return err } + c.activator = act if c.cfg.ProxyTimeout > 0 { - act.SetProxyTimeout(c.cfg.ProxyTimeout) + c.activator.SetProxyTimeout(c.cfg.ProxyTimeout) } if c.cfg.ConnectTimeout > 0 { - act.SetConnectTimeout(c.cfg.ConnectTimeout) + c.activator.SetConnectTimeout(c.cfg.ConnectTimeout) } - c.activator = act } if len(c.cfg.Ports) == 0 { @@ -407,7 +466,7 @@ func (c *Container) initActivator(ctx context.Context, enableRedirects bool) err // ports might fail in various ways. We schedule a retry. retryIn := c.initRetry() log.G(ctx).Infof("no ports detected, retrying init in %s", retryIn) - c.retryInitIn(retryIn, enableRedirects) + c.retryInitIn(retryIn) return nil } @@ -416,16 +475,12 @@ func (c *Container) initActivator(ctx context.Context, enableRedirects bool) err log.G(ctx).Infof("starting activator with ports: %v", c.cfg.Ports) if err := c.startActivator(ctx, c.cfg.Ports...); err != nil { - if errors.Is(err, activator.ErrMapNotFound) { - c.retryInitIn(c.initRetry(), enableRedirects) + if errors.Is(err, activator.ErrMapNotFound) || errors.Is(err, activator.ErrNoListeningSockets) { + c.retryInitIn(c.initRetry()) return nil } return err } - - if enableRedirects { - return c.activator.Reset() - } return nil } @@ -442,10 +497,10 @@ func (c *Container) initRetry() time.Duration { return c.initBackoff } -func (c *Container) retryInitIn(in time.Duration, enableRedirects bool) { +func (c *Container) retryInitIn(in time.Duration) { log.G(c.context).Infof("scheduling init in %s", in) timer := time.AfterFunc(in, func() { - if err := c.initActivator(c.context, enableRedirects); err != nil { + if err := c.initActivator(c.context); err != nil { log.G(c.context).Warnf("error initializing activator: %s", err) } }) @@ -459,6 +514,26 @@ func (c *Container) cancelInit() { c.initTimer.Stop() } +func (c *Container) getListeners(ports ...uint16) activator.Listeners { + lns, err := c.loadListeners() + if err == nil { + return lns + } + + listeners := activator.Listeners{} + for _, port := range ports { + // fallback to just a dual-stack listener for each port. If !skipStart, + // the listeners will anyways be detected from the app so this is only + // relevant if we skipStart and the listeners from the checkpoint are + // empty + listeners = append( + listeners, + activator.Listener{Port: port, Network: activator.NetworkTCPAny}, + ) + } + return listeners +} + // startActivator starts the activator func (c *Container) startActivator(ctx context.Context, ports ...uint16) error { if c.activator.Started() { @@ -468,7 +543,8 @@ func (c *Container) startActivator(ctx context.Context, ports ...uint16) error { log.G(ctx).WithError(err).Error("failed to attach activator") return err } - if err := c.activator.Start(c.context, c.detectProbe(c.context), c.restoreHandler(c.context), ports...); err != nil { + + if err := c.activator.Start(c.context, c.Pid(), c.getListeners(ports...), c.SkipStart()); err != nil { if errors.Is(err, activator.ErrMapNotFound) { return err } @@ -481,18 +557,16 @@ func (c *Container) startActivator(ctx context.Context, ports ...uint16) error { } func (c *Container) restoreHandler(ctx context.Context) activator.RestoreHook { - return func() error { - log.G(ctx).Printf("got a request") - + return func() (int, error) { restoredContainer, _, err := c.Restore(ctx) if err != nil { if errors.Is(err, ErrAlreadyRestored) { log.G(ctx).Info("container is already restored, ignoring request") - return nil + return c.Pid(), nil } if errors.Is(err, ErrNoCapacity) { log.G(ctx).Info("no capacity to restore, requests are being forwarded") - return nil + return 0, nil } // restore failed, this is currently unrecoverable, so we set the // process to exited and let the runtime recreate it. @@ -502,7 +576,7 @@ func (c *Container) restoreHandler(ctx context.Context) activator.RestoreHook { } c.Container = restoredContainer c.ScheduleScaleDown() - return nil + return c.Container.Pid(), nil } } diff --git a/shim/probe.go b/shim/probe.go index 10f15ed..e3905a1 100644 --- a/shim/probe.go +++ b/shim/probe.go @@ -39,14 +39,18 @@ func (c *Container) detectProbe(ctx context.Context) activator.ConnHook { } } -func isKubeletAddr(ctx context.Context, remoteAddr net.Addr, activator *activator.Server) bool { +func isKubeletAddr(ctx context.Context, remoteAddr net.Addr, act activator.Activator) bool { + srv, ok := act.(*activator.Server) + if !ok { + return false + } tcpAddr, ok := remoteAddr.(*net.TCPAddr) if !ok { log.G(ctx).Debugf("remoteAddr is not a *net.TCPAddr: %T", remoteAddr) return false } remoteIPAddr := tcpAddr.AddrPort().Addr().Unmap() - kubeletAddr, err := activator.GetKubeletAddr(remoteIPAddr.Is6()) + kubeletAddr, err := srv.GetKubeletAddr(remoteIPAddr.Is6()) if err != nil { log.G(ctx).WithError(err).Debug("getting kubelet addr") return false diff --git a/shim/restore.go b/shim/restore.go index 1b3b352..971a3f3 100644 --- a/shim/restore.go +++ b/shim/restore.go @@ -48,13 +48,18 @@ func (c *Container) Restore(ctx context.Context) (*runc.Container, process.Proce if !c.ScaledDown() { return nil, nil, ErrAlreadyRestored } + if err := c.reloadConfig(ctx); err != nil { + log.G(ctx).WithError(err).Error("reloading config") + } resp, err := c.restoreCapacityRequest(ctx) if err != nil { // log the error but continue with restoring log.G(ctx).WithError(err).Error("requesting restore capacity") } else if !resp.Allowed { if resp.RedirectAddr != "" { - c.activator.ForwardToTarget(resp.RedirectAddr) + if err := c.activator.ForwardToTarget(ctx, resp.RedirectAddr); err != nil { + return nil, nil, err + } } return nil, nil, ErrNoCapacity } @@ -222,7 +227,7 @@ func createContainerLoggers(ctx context.Context, logPath string, tty bool) (stdo // MigrationRestore requests a restore from the node. If a matching migration is // found, it sets the Checkpoint path in the CreateTaskRequest. -func MigrationRestore(ctx context.Context, r *task.CreateTaskRequest, cfg *v1.Config) (skipStart bool, err error) { +func MigrationRestore(ctx context.Context, r *task.CreateTaskRequest, cfg *v1.Config) (bool, error) { conn, err := net.Dial("unix", nodev1.SocketPath) if err != nil { return false, fmt.Errorf("%w: dialing node service: %w", ErrRestoreDial, err) @@ -274,8 +279,7 @@ func MigrationRestore(ctx context.Context, r *task.CreateTaskRequest, cfg *v1.Co } if !resp.MigrationInfo.LiveMigration { - skipStart = true - return + return true, nil } // wait for the lazy pages socket file to exist to ensure the pages