mirror of
https://github.com/agent-substrate/substrate.git
synced 2026-10-02 03:24:42 +08:00
Clean up packages in ateomnet (#1897)
* Move the sandbox DNS code into its own package * Move the network namespace primitives into their own package
This commit is contained in:
@@ -26,6 +26,7 @@ import (
|
||||
"path/filepath"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/dns"
|
||||
"github.com/agent-substrate/substrate/internal/ateompath"
|
||||
"github.com/agent-substrate/substrate/internal/atunnel"
|
||||
)
|
||||
@@ -40,7 +41,7 @@ func actorResolvConf(actorUID string) (string, error) {
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
|
||||
return "", fmt.Errorf("creating the actor directory: %w", err)
|
||||
}
|
||||
if err := os.WriteFile(path, ateomnet.SandboxResolvConf(pod), 0o644); err != nil {
|
||||
if err := os.WriteFile(path, dns.SandboxResolvConf(ateomnet.ActorVethGateway, pod), 0o644); err != nil {
|
||||
return "", fmt.Errorf("writing the actor resolv.conf: %w", err)
|
||||
}
|
||||
return path, nil
|
||||
|
||||
@@ -23,13 +23,13 @@ import (
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/vishvananda/netns"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/agent-substrate/substrate/cmd/ateom-microvm/internal/kata"
|
||||
"github.com/agent-substrate/substrate/internal/ateomcgroup"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
"github.com/agent-substrate/substrate/internal/atunnel"
|
||||
"github.com/agent-substrate/substrate/internal/resources"
|
||||
"github.com/agent-substrate/substrate/internal/sizing"
|
||||
@@ -218,7 +218,7 @@ func (s *AteomService) guestStatsFor(actorUID string) *guestStatsTarget {
|
||||
}
|
||||
|
||||
// sandboxNetNS is where an actor's tap and atunnel's sockets live, or -1.
|
||||
func (s *AteomService) sandboxNetNS(actorUID string) netns.NsHandle {
|
||||
func (s *AteomService) sandboxNetNS(actorUID string) netns.Handle {
|
||||
hosted := s.lookupActor(actorUID)
|
||||
if hosted == nil {
|
||||
return -1
|
||||
|
||||
@@ -23,9 +23,9 @@ import (
|
||||
"os"
|
||||
|
||||
"github.com/vishvananda/netlink"
|
||||
"github.com/vishvananda/netns"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -51,9 +51,9 @@ var gatewayHWAddr = ateomnet.MustParseMAC(gatewayMAC)
|
||||
|
||||
// setupActorTap creates the guest's tap with a fixed gateway address and MAC.
|
||||
// Returns the FDs cloud-hypervisor adopts on boot or restore.
|
||||
func setupActorTap(ctx context.Context, actorNetNS netns.NsHandle, name string, queuePairs int) ([]*os.File, error) {
|
||||
func setupActorTap(ctx context.Context, actorNetNS netns.Handle, name string, queuePairs int) ([]*os.File, error) {
|
||||
var fds []*os.File
|
||||
err := ateomnet.NetNSDo(ctx, actorNetNS, func(ctx context.Context) error {
|
||||
err := netns.Do(ctx, actorNetNS, func(ctx context.Context) error {
|
||||
if old, lerr := netlink.LinkByName(name); lerr == nil {
|
||||
_ = netlink.LinkDel(old)
|
||||
}
|
||||
@@ -98,9 +98,9 @@ func setupActorTap(ctx context.Context, actorNetNS netns.NsHandle, name string,
|
||||
}
|
||||
|
||||
// actorTapMTUOf reads the tap MTU, falling back to actorTapMTU on error.
|
||||
func actorTapMTUOf(ctx context.Context, actorNetNS netns.NsHandle, name string) int {
|
||||
func actorTapMTUOf(ctx context.Context, actorNetNS netns.Handle, name string) int {
|
||||
mtu := actorTapMTU
|
||||
_ = ateomnet.NetNSDo(ctx, actorNetNS, func(ctx context.Context) error {
|
||||
_ = netns.Do(ctx, actorNetNS, func(ctx context.Context) error {
|
||||
if l, err := netlink.LinkByName(name); err == nil {
|
||||
mtu = l.Attrs().MTU
|
||||
} else {
|
||||
|
||||
@@ -21,9 +21,9 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/vishvananda/netlink"
|
||||
"github.com/vishvananda/netns"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
"github.com/agent-substrate/substrate/internal/atunnel"
|
||||
"github.com/agent-substrate/substrate/internal/nodepath"
|
||||
"github.com/agent-substrate/substrate/internal/resources"
|
||||
@@ -72,15 +72,15 @@ func TestHostActorReplacesSameActor(t *testing.T) {
|
||||
}
|
||||
|
||||
// tapNetNS gives a test its own namespace to build a tap in.
|
||||
func tapNetNS(t *testing.T, name string) netns.NsHandle {
|
||||
func tapNetNS(t *testing.T, name string) netns.Handle {
|
||||
t.Helper()
|
||||
ns, err := ateomnet.CreateNetNSWithoutSwitching(name)
|
||||
ns, err := netns.CreateNamed(name)
|
||||
if err != nil {
|
||||
t.Fatalf("creating namespace: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ns.Close()
|
||||
_ = netns.DeleteNamed(name)
|
||||
_ = netns.RemoveNamed(name)
|
||||
})
|
||||
return ns
|
||||
}
|
||||
@@ -106,7 +106,7 @@ func TestSetupActorTap(t *testing.T) {
|
||||
t.Errorf("got %d descriptors, want one per queue pair", len(fds))
|
||||
}
|
||||
|
||||
if err := ateomnet.NetNSDo(ctx, ns, func(context.Context) error {
|
||||
if err := netns.Do(ctx, ns, func(context.Context) error {
|
||||
link, err := netlink.LinkByName("tap0_kata")
|
||||
if err != nil {
|
||||
return err
|
||||
|
||||
@@ -21,6 +21,7 @@ import (
|
||||
"os"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/dns"
|
||||
"github.com/agent-substrate/substrate/internal/atunnel"
|
||||
)
|
||||
|
||||
@@ -30,7 +31,7 @@ func writeActorResolvConf(rootfs string) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("reading the worker pod resolv.conf: %w", err)
|
||||
}
|
||||
return ateomnet.WriteRootfsResolvConf(rootfs, ateomnet.SandboxResolvConf(pod))
|
||||
return dns.WriteRootfsResolvConf(rootfs, dns.SandboxResolvConf(ateomnet.ActorVethGateway, pod))
|
||||
}
|
||||
|
||||
// attachAtunnel completes setup after atunnel receives the service's dialer.
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
@@ -14,7 +12,9 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package ateomnet
|
||||
// Package dns answers an actor's DNS from its sandbox's gateway namespace and
|
||||
// writes the resolv.conf that points the actor at it.
|
||||
package dns
|
||||
|
||||
import (
|
||||
"errors"
|
||||
@@ -25,11 +25,12 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// SandboxResolvConf replaces nameservers with the sandbox gateway while
|
||||
// preserving the pod's search domains and options for Kubernetes DNS.
|
||||
func SandboxResolvConf(podResolvConf []byte) []byte {
|
||||
// SandboxResolvConf replaces the pod's nameservers with nameserver, the
|
||||
// address the sandbox's DNS is served on, while preserving the pod's search
|
||||
// domains and options for Kubernetes DNS.
|
||||
func SandboxResolvConf(nameserver string, podResolvConf []byte) []byte {
|
||||
var out strings.Builder
|
||||
out.WriteString("nameserver " + ActorVethGateway + "\n")
|
||||
out.WriteString("nameserver " + nameserver + "\n")
|
||||
for line := range strings.SplitSeq(string(podResolvConf), "\n") {
|
||||
if strings.HasPrefix(strings.TrimSpace(line), "nameserver") {
|
||||
continue
|
||||
@@ -47,7 +48,7 @@ func SandboxResolvConf(podResolvConf []byte) []byte {
|
||||
// os.Root confines path traversal; unlinking prevents writes through existing links.
|
||||
func WriteRootfsResolvConf(rootfs string, content []byte) error {
|
||||
if len(content) == 0 {
|
||||
return fmt.Errorf("actornet: refusing to write an empty resolv.conf")
|
||||
return fmt.Errorf("dns: refusing to write an empty resolv.conf")
|
||||
}
|
||||
root, err := os.OpenRoot(rootfs)
|
||||
if err != nil {
|
||||
+3
-5
@@ -1,5 +1,3 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
@@ -14,7 +12,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package ateomnet
|
||||
package dns
|
||||
|
||||
import (
|
||||
"os"
|
||||
@@ -31,7 +29,7 @@ func TestSandboxResolvConf(t *testing.T) {
|
||||
"search ate-demo.svc.cluster.local svc.cluster.local cluster.local\n" +
|
||||
"options ndots:5\n"
|
||||
|
||||
got := string(SandboxResolvConf([]byte(pod)))
|
||||
got := string(SandboxResolvConf("169.254.17.1", []byte(pod)))
|
||||
|
||||
want := "nameserver 169.254.17.1\n" +
|
||||
"search ate-demo.svc.cluster.local svc.cluster.local cluster.local\n" +
|
||||
@@ -45,7 +43,7 @@ func TestSandboxResolvConf(t *testing.T) {
|
||||
// only a nameserver line resolves public names but not cluster ones.
|
||||
func TestSandboxResolvConfKeepsSearchAndOptions(t *testing.T) {
|
||||
pod := "search svc.cluster.local\nnameserver 10.96.0.10\nnameserver 10.96.0.11\noptions ndots:5 timeout:1\n"
|
||||
got := string(SandboxResolvConf([]byte(pod)))
|
||||
got := string(SandboxResolvConf("169.254.17.1", []byte(pod)))
|
||||
|
||||
if strings.Contains(got, "10.96.0.10") || strings.Contains(got, "10.96.0.11") {
|
||||
t.Errorf("a pod resolver survived into the actor's file, so its DNS would bypass atunnel:\n%s", got)
|
||||
@@ -14,7 +14,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package ateomnet
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -24,24 +24,24 @@ import (
|
||||
"net"
|
||||
"strconv"
|
||||
|
||||
"github.com/vishvananda/netns"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
)
|
||||
|
||||
// dnsServer answers an actor's DNS. Satisfied by atunnel.DNSRelay; an interface
|
||||
// Server answers an actor's DNS. Satisfied by atunnel.DNSRelay; an interface
|
||||
// so this package does not depend on it.
|
||||
type dnsServer interface {
|
||||
type Server interface {
|
||||
ServePacket(ctx context.Context, pc net.PacketConn) error
|
||||
Serve(ctx context.Context, listener net.Listener) error
|
||||
}
|
||||
|
||||
// serveSandboxDNS serves UDP and TCP DNS in the sandbox's local gateway namespace.
|
||||
func serveSandboxDNS(ctx context.Context, relay dnsServer, ns netns.NsHandle, port uint16) (_ []io.Closer, _ []func(), retErr error) {
|
||||
// Serve serves UDP and TCP DNS in the sandbox's local gateway namespace.
|
||||
func Serve(ctx context.Context, relay Server, ns netns.Handle, port uint16) ([]io.Closer, []func(), error) {
|
||||
// Bind the wildcard because the microVM tap's gateway address is added later.
|
||||
address := net.JoinHostPort("0.0.0.0", strconv.Itoa(int(port)))
|
||||
|
||||
var packet net.PacketConn
|
||||
var stream net.Listener
|
||||
if err := NetNSDo(ctx, ns, func(context.Context) error {
|
||||
if err := netns.Do(ctx, ns, func(context.Context) error {
|
||||
pc, err := net.ListenPacket("udp", address)
|
||||
if err != nil {
|
||||
return fmt.Errorf("while opening the actor DNS socket: %w", err)
|
||||
@@ -78,3 +78,9 @@ func serveSandboxDNS(ctx context.Context, relay dnsServer, ns netns.NsHandle, po
|
||||
closers := []io.Closer{closerFunc(func() error { stopServing(); return nil }), packet, stream}
|
||||
return closers, serve, nil
|
||||
}
|
||||
|
||||
// closerFunc adapts a cancel function to io.Closer, so a caller takes a
|
||||
// sandbox's sockets and the work behind them down as one list.
|
||||
type closerFunc func() error
|
||||
|
||||
func (f closerFunc) Close() error { return f() }
|
||||
@@ -0,0 +1,78 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
"github.com/agent-substrate/substrate/internal/roottest"
|
||||
)
|
||||
|
||||
// stoppableDNS records that its serving contexts were canceled.
|
||||
type stoppableDNS struct{ packet, stream chan struct{} }
|
||||
|
||||
func (d *stoppableDNS) ServePacket(ctx context.Context, pc net.PacketConn) error {
|
||||
<-ctx.Done()
|
||||
close(d.packet)
|
||||
return pc.Close()
|
||||
}
|
||||
|
||||
func (d *stoppableDNS) Serve(ctx context.Context, l net.Listener) error {
|
||||
<-ctx.Done()
|
||||
close(d.stream)
|
||||
return l.Close()
|
||||
}
|
||||
|
||||
func TestClosingSandboxDNSStopsServing(t *testing.T) {
|
||||
roottest.Require(t, "creates network namespaces")
|
||||
const nsName = "dns-teardown-test"
|
||||
ns, err := netns.CreateNamed(nsName)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() {
|
||||
ns.Close()
|
||||
_ = netns.RemoveNamed(nsName)
|
||||
}()
|
||||
|
||||
relay := &stoppableDNS{packet: make(chan struct{}), stream: make(chan struct{})}
|
||||
closers, serve, err := Serve(context.Background(), relay, ns, 53)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, fn := range serve {
|
||||
go fn()
|
||||
}
|
||||
for _, c := range closers {
|
||||
_ = c.Close()
|
||||
}
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
stopped chan struct{}
|
||||
}{{"UDP", relay.packet}, {"TCP", relay.stream}} {
|
||||
select {
|
||||
case <-tc.stopped:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Errorf("%s serving outlived the sandbox's sockets", tc.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -18,19 +18,11 @@
|
||||
package ateomnet
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
"github.com/google/nftables/expr"
|
||||
"github.com/vishvananda/netlink"
|
||||
"github.com/vishvananda/netns"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -80,40 +72,6 @@ func MustParseMAC(s string) net.HardwareAddr {
|
||||
return m
|
||||
}
|
||||
|
||||
// AllowUnprivilegedPorts lets this namespace bind ports below 1024 without
|
||||
// CAP_NET_BIND_SERVICE, which is how atunnel answers a sandbox's DNS on 53.
|
||||
// The sysctl is per-namespace and grants nothing outside it.
|
||||
func AllowUnprivilegedPorts() error {
|
||||
return setNetSysctl("net/ipv4/ip_unprivileged_port_start", "0")
|
||||
}
|
||||
|
||||
// setNetSysctl writes value to the named sysctl in the current network
|
||||
// namespace, remounting /proc/sys read-write when the runtime bind-mounted it
|
||||
// read-only. A no-op when it already reads that way.
|
||||
func setNetSysctl(key, value string) error {
|
||||
path := "/proc/sys/" + key
|
||||
if b, err := os.ReadFile(path); err == nil && strings.TrimSpace(string(b)) == value {
|
||||
return nil
|
||||
}
|
||||
// Only EROFS is worth remounting for; any other error is returned as is.
|
||||
if err := os.WriteFile(path, []byte(value+"\n"), 0o644); !errors.Is(err, unix.EROFS) {
|
||||
if err != nil {
|
||||
return fmt.Errorf("while setting %s in worker pod netns: %w", key, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err := unix.Mount("none", "/proc/sys", "", unix.MS_BIND|unix.MS_REMOUNT, ""); err != nil {
|
||||
return fmt.Errorf("while remounting /proc/sys read-write to set %s: %w", key, err)
|
||||
}
|
||||
defer func() {
|
||||
_ = unix.Mount("none", "/proc/sys", "", unix.MS_BIND|unix.MS_REMOUNT|unix.MS_RDONLY, "")
|
||||
}()
|
||||
if err := os.WriteFile(path, []byte(value+"\n"), 0o644); err != nil {
|
||||
return fmt.Errorf("while setting %s in worker pod netns: %w", key, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func l4ProtocolEqual(proto byte) []expr.Any {
|
||||
return []expr.Any{
|
||||
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
|
||||
@@ -124,85 +82,3 @@ func l4ProtocolEqual(proto byte) []expr.Any {
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// CreateNetNSWithoutSwitching creates a named netns and returns its handle,
|
||||
// restoring the caller's current netns before returning.
|
||||
//
|
||||
// The caller owns the name exclusively, so a name still present when this
|
||||
// runs was left behind by an earlier incarnation and is removed first. The
|
||||
// kernel creates the name with O_EXCL, so without that removal a single
|
||||
// failed teardown would wedge the name for good: nothing could ever create
|
||||
// it again. Removal only unmounts and unlinks the name. Anything still
|
||||
// holding the namespace keeps it alive, and existing handles stay usable.
|
||||
func CreateNetNSWithoutSwitching(name string) (netns.NsHandle, error) {
|
||||
runtime.LockOSThread()
|
||||
defer runtime.UnlockOSThread()
|
||||
|
||||
if err := removeNamedNetNS(name); err != nil {
|
||||
return -1, fmt.Errorf("while removing the leftover netns %s: %w", name, err)
|
||||
}
|
||||
|
||||
// We need to create the new NS, then switch back to the current netns.
|
||||
curNetNS, err := netns.Get()
|
||||
if err != nil {
|
||||
return -1, fmt.Errorf("while getting current netns: %w", err)
|
||||
}
|
||||
// Registered before the restoring defer below since deferred calls are LIFO.
|
||||
defer curNetNS.Close()
|
||||
defer func() {
|
||||
if err := netns.Set(curNetNS); err != nil {
|
||||
// Better to blow up the program than continue execution with
|
||||
// one OS thread randomly in a different netns.
|
||||
panic(fmt.Sprintf("Failed to restore original netns: %v", err))
|
||||
}
|
||||
}()
|
||||
|
||||
interiorNetNS, err := netns.NewNamed(name)
|
||||
if err != nil {
|
||||
return -1, fmt.Errorf("while creating interior network namespace: %w", err)
|
||||
}
|
||||
return interiorNetNS, nil
|
||||
}
|
||||
|
||||
func removeNamedNetNS(name string) error {
|
||||
if name == "" || name == "." || name == ".." || strings.ContainsAny(name, "/\x00") {
|
||||
return fmt.Errorf("invalid network namespace name %q: %w", name, os.ErrInvalid)
|
||||
}
|
||||
path := filepath.Join("/run/netns", name)
|
||||
if err := unix.Unmount(path, unix.MNT_DETACH|unix.UMOUNT_NOFOLLOW); err != nil && !errors.Is(err, unix.ENOENT) && !errors.Is(err, unix.EINVAL) {
|
||||
return err
|
||||
}
|
||||
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// NetNSDo runs do() with the OS thread switched into targetNS, then restores it.
|
||||
func NetNSDo(ctx context.Context, targetNS netns.NsHandle, do func(context.Context) error) error {
|
||||
runtime.LockOSThread()
|
||||
defer runtime.UnlockOSThread()
|
||||
|
||||
// We need to create the new NS, then switch back to the current netns.
|
||||
curNetNS, err := netns.Get()
|
||||
if err != nil {
|
||||
return fmt.Errorf("while getting current netns: %w", err)
|
||||
}
|
||||
// Registered before the restoring defer below since deferred calls are LIFO.
|
||||
defer curNetNS.Close()
|
||||
defer func() {
|
||||
if err := netns.Set(curNetNS); err != nil {
|
||||
// Better to blow up the program than continue execution with
|
||||
// one OS thread randomly in a different netns.
|
||||
panic(fmt.Sprintf("Failed to restore original netns: %v", err))
|
||||
}
|
||||
}()
|
||||
|
||||
if err := netns.Set(targetNS); err != nil {
|
||||
return fmt.Errorf("setting target netns: %w", err)
|
||||
}
|
||||
if err := do(ctx); err != nil {
|
||||
return fmt.Errorf("while executing function in target netns: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package netns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"runtime"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
vishnetns "github.com/vishvananda/netns"
|
||||
)
|
||||
|
||||
// Dialer dials TCP or UDP IP literals in ns, pinning a thread only until
|
||||
// the socket is created.
|
||||
func Dialer(ns Handle) func(context.Context, string, string) (net.Conn, error) {
|
||||
return func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateDialTarget(network, addr); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var conn net.Conn
|
||||
dialErr := with(ns, func(restore func() error) error {
|
||||
// Only creating the socket needs the namespace, and
|
||||
// ControlContext runs once it exists: restore there rather than
|
||||
// holding the thread for the whole connect.
|
||||
socketCreated := false
|
||||
dialer := net.Dialer{ControlContext: func(context.Context, string, string, syscall.RawConn) error {
|
||||
if socketCreated {
|
||||
return errors.New("sandbox dial cannot recreate its socket outside the namespace")
|
||||
}
|
||||
socketCreated = true
|
||||
return restore()
|
||||
}}
|
||||
var err error
|
||||
conn, err = dialer.DialContext(ctx, network, addr)
|
||||
return err
|
||||
})
|
||||
if dialErr != nil || ctx.Err() != nil {
|
||||
if conn != nil {
|
||||
_ = conn.Close()
|
||||
}
|
||||
if dialErr != nil {
|
||||
return nil, dialErr
|
||||
}
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
}
|
||||
|
||||
func validateDialTarget(network, addr string) error {
|
||||
switch network {
|
||||
case "tcp", "tcp4", "tcp6", "udp", "udp4", "udp6":
|
||||
default:
|
||||
return net.UnknownNetworkError(network)
|
||||
}
|
||||
hostname, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := netip.ParseAddr(hostname); err != nil {
|
||||
return fmt.Errorf("netns.Dialer supports only IP literals (got %q): %w", hostname, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// with switches to targetNS, calls run, then restores the original namespace.
|
||||
// run can call restore to switch back and unlock the OS thread before returning.
|
||||
// Calling restore again after it succeeds has no effect.
|
||||
//
|
||||
// A separate goroutine lets us leave the thread locked if restoration fails.
|
||||
// Go then discards that thread when the goroutine exits.
|
||||
func with(targetNS Handle, run func(restore func() error) error) error {
|
||||
var resultErr error
|
||||
var done sync.WaitGroup
|
||||
done.Add(1)
|
||||
go func() {
|
||||
defer done.Done()
|
||||
runtime.LockOSThread()
|
||||
originalNS, err := vishnetns.Get()
|
||||
if err != nil {
|
||||
runtime.UnlockOSThread()
|
||||
resultErr = fmt.Errorf("while reading the current netns: %w", err)
|
||||
return
|
||||
}
|
||||
defer originalNS.Close()
|
||||
if err := vishnetns.Set(targetNS); err != nil {
|
||||
runtime.UnlockOSThread()
|
||||
resultErr = fmt.Errorf("while entering the actor netns: %w", err)
|
||||
return
|
||||
}
|
||||
|
||||
restored := false
|
||||
restore := func() error {
|
||||
if restored {
|
||||
return nil
|
||||
}
|
||||
if err := vishnetns.Set(originalNS); err != nil {
|
||||
return fmt.Errorf("while restoring the worker netns: %w", err)
|
||||
}
|
||||
runtime.UnlockOSThread()
|
||||
restored = true
|
||||
return nil
|
||||
}
|
||||
|
||||
resultErr = run(restore)
|
||||
if err := restore(); err != nil {
|
||||
resultErr = err
|
||||
}
|
||||
}()
|
||||
done.Wait()
|
||||
return resultErr
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package netns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateDialTarget(t *testing.T) {
|
||||
for _, target := range []struct {
|
||||
network, address string
|
||||
wantErr bool
|
||||
}{
|
||||
{"tcp", "127.0.0.1:80", false},
|
||||
{"tcp4", "127.0.0.1:80", false},
|
||||
{"tcp6", "[::1]:80", false},
|
||||
{"udp", "127.0.0.1:53", false},
|
||||
{"udp4", "127.0.0.1:53", false},
|
||||
{"udp6", "[fe80::1%eth0]:53", false},
|
||||
{"tcp", "localhost:80", true},
|
||||
{"tcp", ":80", true},
|
||||
{"tcp", "127.0.0.1", true},
|
||||
{"unix", "/tmp/socket", true},
|
||||
{"ip", "127.0.0.1:80", true},
|
||||
} {
|
||||
t.Run(target.network+"/"+target.address, func(t *testing.T) {
|
||||
err := validateDialTarget(target.network, target.address)
|
||||
if (err != nil) != target.wantErr {
|
||||
t.Errorf("validateDialTarget = %v, want error: %t", err, target.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
if err := validateDialTarget("unix", "/tmp/socket"); !errors.Is(err, net.UnknownNetworkError("unix")) {
|
||||
t.Errorf("unsupported network: got %v, want UnknownNetworkError", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDialerRejectsNonIPTargets(t *testing.T) {
|
||||
for _, target := range []struct{ network, address string }{
|
||||
{"tcp", "localhost:80"},
|
||||
{"tcp", ":80"},
|
||||
{"tcp", "127.0.0.1"},
|
||||
{"unix", "/tmp/socket"},
|
||||
} {
|
||||
if conn, err := Dialer(-1)(context.Background(), target.network, target.address); err == nil {
|
||||
_ = conn.Close()
|
||||
t.Errorf("accepted %s %s", target.network, target.address)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDialerCanceledContext(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if _, err := Dialer(-1)(ctx, "tcp", "127.0.0.1:1"); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("canceled dial: got %v, want cancellation", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
// Package netns creates Linux network namespaces and runs code, opens sockets
|
||||
// and dials inside them. Entering a namespace is a property of the OS thread,
|
||||
// so everything here locks a thread for as long as it is in one and restores
|
||||
// the caller's namespace before returning.
|
||||
package netns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
vishnetns "github.com/vishvananda/netns"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// Handle is a descriptor for a network namespace. Callers name it through this
|
||||
// package so that nothing else has to import the one underneath.
|
||||
type Handle = vishnetns.NsHandle
|
||||
|
||||
// GetFromName opens the namespace of that name under /run/netns.
|
||||
func GetFromName(name string) (Handle, error) {
|
||||
return vishnetns.GetFromName(name)
|
||||
}
|
||||
|
||||
// CreateNamed creates a named netns and returns its handle, restoring the
|
||||
// caller's current netns before returning.
|
||||
//
|
||||
// The caller owns the name exclusively, so a name still present when this
|
||||
// runs was left behind by an earlier incarnation and is removed first. The
|
||||
// kernel creates the name with O_EXCL, so without that removal a single
|
||||
// failed teardown would wedge the name for good: nothing could ever create
|
||||
// it again. Removal only unmounts and unlinks the name. Anything still
|
||||
// holding the namespace keeps it alive, and existing handles stay usable.
|
||||
func CreateNamed(name string) (Handle, error) {
|
||||
runtime.LockOSThread()
|
||||
defer runtime.UnlockOSThread()
|
||||
|
||||
if err := RemoveNamed(name); err != nil {
|
||||
return -1, fmt.Errorf("while removing the leftover netns %s: %w", name, err)
|
||||
}
|
||||
|
||||
// We need to create the new NS, then switch back to the current netns.
|
||||
curNetNS, err := vishnetns.Get()
|
||||
if err != nil {
|
||||
return -1, fmt.Errorf("while getting current netns: %w", err)
|
||||
}
|
||||
// Registered before the restoring defer below since deferred calls are LIFO.
|
||||
defer curNetNS.Close()
|
||||
defer func() {
|
||||
if err := vishnetns.Set(curNetNS); err != nil {
|
||||
// Better to blow up the program than continue execution with
|
||||
// one OS thread randomly in a different netns.
|
||||
panic(fmt.Sprintf("Failed to restore original netns: %v", err))
|
||||
}
|
||||
}()
|
||||
|
||||
interiorNetNS, err := vishnetns.NewNamed(name)
|
||||
if err != nil {
|
||||
return -1, fmt.Errorf("while creating interior network namespace: %w", err)
|
||||
}
|
||||
return interiorNetNS, nil
|
||||
}
|
||||
|
||||
// RemoveNamed unmounts and unlinks a name under /run/netns, without following
|
||||
// it if it is a symlink. A name that is already gone is not an error, and the
|
||||
// namespace itself survives for as long as something holds it open.
|
||||
func RemoveNamed(name string) error {
|
||||
if name == "" || name == "." || name == ".." || strings.ContainsAny(name, "/\x00") {
|
||||
return fmt.Errorf("invalid network namespace name %q: %w", name, os.ErrInvalid)
|
||||
}
|
||||
path := filepath.Join("/run/netns", name)
|
||||
if err := unix.Unmount(path, unix.MNT_DETACH|unix.UMOUNT_NOFOLLOW); err != nil && !errors.Is(err, unix.ENOENT) && !errors.Is(err, unix.EINVAL) {
|
||||
return err
|
||||
}
|
||||
if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Do runs do() with the OS thread switched into targetNS, then restores it.
|
||||
func Do(ctx context.Context, targetNS Handle, do func(context.Context) error) error {
|
||||
runtime.LockOSThread()
|
||||
defer runtime.UnlockOSThread()
|
||||
|
||||
// We need to create the new NS, then switch back to the current netns.
|
||||
curNetNS, err := vishnetns.Get()
|
||||
if err != nil {
|
||||
return fmt.Errorf("while getting current netns: %w", err)
|
||||
}
|
||||
// Registered before the restoring defer below since deferred calls are LIFO.
|
||||
defer curNetNS.Close()
|
||||
defer func() {
|
||||
if err := vishnetns.Set(curNetNS); err != nil {
|
||||
// Better to blow up the program than continue execution with
|
||||
// one OS thread randomly in a different netns.
|
||||
panic(fmt.Sprintf("Failed to restore original netns: %v", err))
|
||||
}
|
||||
}()
|
||||
|
||||
if err := vishnetns.Set(targetNS); err != nil {
|
||||
return fmt.Errorf("setting target netns: %w", err)
|
||||
}
|
||||
if err := do(ctx); err != nil {
|
||||
return fmt.Errorf("while executing function in target netns: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Listen opens wildcard TCP listeners inside ns.
|
||||
// Sockets retain their namespace and can be served from another namespace.
|
||||
func Listen(ctx context.Context, ns Handle, ports []uint16) (_ []net.Listener, retErr error) {
|
||||
var listeners []net.Listener
|
||||
defer func() {
|
||||
if retErr != nil {
|
||||
for _, l := range listeners {
|
||||
_ = l.Close()
|
||||
}
|
||||
}
|
||||
}()
|
||||
if err := Do(ctx, ns, func(context.Context) error {
|
||||
for _, port := range ports {
|
||||
l, err := net.Listen("tcp", fmt.Sprintf("0.0.0.0:%d", port))
|
||||
if err != nil {
|
||||
return fmt.Errorf("while listening on port %d: %w", port, err)
|
||||
}
|
||||
listeners = append(listeners, l)
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return listeners, nil
|
||||
}
|
||||
@@ -14,7 +14,7 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package ateomnet
|
||||
package netns
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -30,32 +30,32 @@ import (
|
||||
func TestNamedNetNSRejectsInvalidNames(t *testing.T) {
|
||||
for _, name := range []string{"", ".", "..", "/absolute", "../outside", "nested/name", "ateom-actor:uid/../../outside", "nul\x00name"} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if err := removeNamedNetNS(name); !errors.Is(err, os.ErrInvalid) {
|
||||
t.Fatalf("removeNamedNetNS(%q): got %v, want invalid name", name, err)
|
||||
if err := RemoveNamed(name); !errors.Is(err, os.ErrInvalid) {
|
||||
t.Fatalf("RemoveNamed(%q): got %v, want invalid name", name, err)
|
||||
}
|
||||
handle, err := CreateNetNSWithoutSwitching(name)
|
||||
handle, err := CreateNamed(name)
|
||||
if err == nil {
|
||||
handle.Close()
|
||||
}
|
||||
if !errors.Is(err, os.ErrInvalid) {
|
||||
t.Fatalf("CreateNetNSWithoutSwitching(%q): got %v, want invalid name", name, err)
|
||||
t.Fatalf("CreateNamed(%q): got %v, want invalid name", name, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveNamedNetNSDoesNotFollowSymlinks(t *testing.T) {
|
||||
func TestRemoveNamedDoesNotFollowSymlinks(t *testing.T) {
|
||||
roottest.Require(t, "creates network namespaces")
|
||||
const targetName = "ateomnet-symlink-target-test"
|
||||
const linkName = "ateomnet-symlink-test"
|
||||
targetPath := "/run/netns/" + targetName
|
||||
linkPath := "/run/netns/" + linkName
|
||||
target, err := CreateNetNSWithoutSwitching(targetName)
|
||||
target, err := CreateNamed(targetName)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer target.Close()
|
||||
t.Cleanup(func() { _ = removeNamedNetNS(targetName) })
|
||||
t.Cleanup(func() { _ = RemoveNamed(targetName) })
|
||||
before, err := os.Stat(targetPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -64,7 +64,7 @@ func TestRemoveNamedNetNSDoesNotFollowSymlinks(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.Remove(linkPath) })
|
||||
if err := removeNamedNetNS(linkName); err != nil {
|
||||
if err := RemoveNamed(linkName); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := os.Lstat(linkPath); !errors.Is(err, os.ErrNotExist) {
|
||||
@@ -79,19 +79,19 @@ func TestRemoveNamedNetNSDoesNotFollowSymlinks(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateNetNSWithoutSwitchingReplacesALeftover(t *testing.T) {
|
||||
func TestCreateNamedReplacesALeftover(t *testing.T) {
|
||||
roottest.Require(t, "creates network namespaces")
|
||||
for _, state := range []string{"mounted", "unmounted"} {
|
||||
t.Run(state, func(t *testing.T) {
|
||||
name := "ateomnet-leftover-test-" + state
|
||||
path := "/run/netns/" + name
|
||||
t.Cleanup(func() { _ = removeNamedNetNS(name) })
|
||||
t.Cleanup(func() { _ = RemoveNamed(name) })
|
||||
|
||||
// Held open across the replacement below: unlinking the name
|
||||
// must not invalidate a handle the caller still has.
|
||||
first, err := CreateNetNSWithoutSwitching(name)
|
||||
first, err := CreateNamed(name)
|
||||
if err != nil {
|
||||
t.Fatalf("first CreateNetNSWithoutSwitching: %v", err)
|
||||
t.Fatalf("first CreateNamed: %v", err)
|
||||
}
|
||||
defer first.Close()
|
||||
if state == "unmounted" {
|
||||
@@ -103,7 +103,7 @@ func TestCreateNetNSWithoutSwitchingReplacesALeftover(t *testing.T) {
|
||||
t.Fatalf("expected the leftover netns to remain: %v", err)
|
||||
}
|
||||
|
||||
second, err := CreateNetNSWithoutSwitching(name)
|
||||
second, err := CreateNamed(name)
|
||||
if err != nil {
|
||||
t.Fatalf("the name is wedged by its own leftover: %v", err)
|
||||
}
|
||||
@@ -119,14 +119,14 @@ func TestCreateNetNSWithoutSwitchingReplacesALeftover(t *testing.T) {
|
||||
if !first.IsOpen() {
|
||||
t.Error("the retained handle closed when its name was replaced")
|
||||
}
|
||||
if err := NetNSDo(context.Background(), first, func(context.Context) error {
|
||||
if err := Do(context.Background(), first, func(context.Context) error {
|
||||
_, err := netlink.LinkList()
|
||||
return err
|
||||
}); err != nil {
|
||||
t.Errorf("the retained handle is no longer usable: %v", err)
|
||||
}
|
||||
for range 2 {
|
||||
if err := removeNamedNetNS(name); err != nil {
|
||||
if err := RemoveNamed(name); err != nil {
|
||||
t.Fatalf("removing namespace: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) {
|
||||
@@ -136,20 +136,3 @@ func TestCreateNetNSWithoutSwitchingReplacesALeftover(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Only EROFS takes the remount path: any other error is reported as it is,
|
||||
// and /proc/sys is left as it was found. Remounting it read-only on the way
|
||||
// out would break every later write.
|
||||
func TestSetNetSysctlReportsAnUnrelatedError(t *testing.T) {
|
||||
err := setNetSysctl("net/ipv4/ateomnet_no_such_sysctl", "0")
|
||||
if !errors.Is(err, unix.ENOENT) {
|
||||
t.Fatalf("setNetSysctl() on a missing key: got %v, want ENOENT", err)
|
||||
}
|
||||
var st unix.Statfs_t
|
||||
if err := unix.Statfs("/proc/sys", &st); err != nil {
|
||||
t.Fatalf("statfs /proc/sys: %v", err)
|
||||
}
|
||||
if st.Flags&unix.ST_RDONLY != 0 {
|
||||
t.Error("/proc/sys was left read-only")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package netns
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// AllowUnprivilegedPorts lets this namespace bind ports below 1024 without
|
||||
// CAP_NET_BIND_SERVICE, which is how atunnel answers a sandbox's DNS on 53.
|
||||
// The sysctl is per-namespace and grants nothing outside it.
|
||||
func AllowUnprivilegedPorts() error {
|
||||
return setSysctl("net/ipv4/ip_unprivileged_port_start", "0")
|
||||
}
|
||||
|
||||
// setSysctl writes value to the named sysctl in the current network
|
||||
// namespace, remounting /proc/sys read-write when the runtime bind-mounted it
|
||||
// read-only. A no-op when it already reads that way.
|
||||
func setSysctl(key, value string) error {
|
||||
path := "/proc/sys/" + key
|
||||
if b, err := os.ReadFile(path); err == nil && strings.TrimSpace(string(b)) == value {
|
||||
return nil
|
||||
}
|
||||
// Only EROFS is worth remounting for; any other error is returned as is.
|
||||
if err := os.WriteFile(path, []byte(value+"\n"), 0o644); !errors.Is(err, unix.EROFS) {
|
||||
if err != nil {
|
||||
return fmt.Errorf("while setting %s in the current netns: %w", key, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err := unix.Mount("none", "/proc/sys", "", unix.MS_BIND|unix.MS_REMOUNT, ""); err != nil {
|
||||
return fmt.Errorf("while remounting /proc/sys read-write to set %s: %w", key, err)
|
||||
}
|
||||
defer func() {
|
||||
_ = unix.Mount("none", "/proc/sys", "", unix.MS_BIND|unix.MS_REMOUNT|unix.MS_RDONLY, "")
|
||||
}()
|
||||
if err := os.WriteFile(path, []byte(value+"\n"), 0o644); err != nil {
|
||||
return fmt.Errorf("while setting %s in the current netns: %w", key, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
//go:build linux
|
||||
|
||||
// Copyright 2026 Google LLC
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package netns
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
// Only EROFS takes the remount path: any other error is reported as it is,
|
||||
// and /proc/sys is left as it was found. Remounting it read-only on the way
|
||||
// out would break every later write.
|
||||
func TestSetSysctlReportsAnUnrelatedError(t *testing.T) {
|
||||
err := setSysctl("net/ipv4/ateomnet_no_such_sysctl", "0")
|
||||
if !errors.Is(err, unix.ENOENT) {
|
||||
t.Fatalf("setSysctl() on a missing key: got %v, want ENOENT", err)
|
||||
}
|
||||
var st unix.Statfs_t
|
||||
if err := unix.Statfs("/proc/sys", &st); err != nil {
|
||||
t.Fatalf("statfs /proc/sys: %v", err)
|
||||
}
|
||||
if st.Flags&unix.ST_RDONLY != 0 {
|
||||
t.Error("/proc/sys was left read-only")
|
||||
}
|
||||
}
|
||||
+24
-163
@@ -23,18 +23,16 @@ import (
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"runtime"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/dns"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
"github.com/agent-substrate/substrate/internal/nodepath"
|
||||
|
||||
"github.com/google/nftables"
|
||||
"github.com/google/nftables/binaryutil"
|
||||
"github.com/google/nftables/expr"
|
||||
"github.com/vishvananda/netlink"
|
||||
"github.com/vishvananda/netns"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
@@ -46,7 +44,7 @@ type SandboxNetwork struct {
|
||||
// RuntimeNetNS is what the sandbox runs in, whichever runtime that is: the
|
||||
// micro-VM's tap lives here, and gVisor claims every interface here and
|
||||
// moves their addresses into its own stack.
|
||||
RuntimeNetNS netns.NsHandle
|
||||
RuntimeNetNS netns.Handle
|
||||
// GatewayNetNS holds the sandbox's default gateway, DNS relay, and atunnel
|
||||
// sockets. This is local to the sandbox, not the external egress gateway.
|
||||
// For microVMs it shares RuntimeNetNS.
|
||||
@@ -54,7 +52,7 @@ type SandboxNetwork struct {
|
||||
// TODO: we hope gVisor can take that same single-namespace shape soon,
|
||||
// once runsc can be given one interface rather than claiming every
|
||||
// interface in the namespace it runs in.
|
||||
GatewayNetNS netns.NsHandle
|
||||
GatewayNetNS netns.Handle
|
||||
}
|
||||
|
||||
func (n *SandboxNetwork) holdsNetNS() bool { return n.RuntimeNetNS > 0 }
|
||||
@@ -88,14 +86,14 @@ func SetupSandboxNetwork(ctx context.Context, cfg SandboxNetworkConfig) (_ *Sand
|
||||
}
|
||||
|
||||
actorNSName := nodepath.ActorNetNSName(actorUID)
|
||||
actorNS, err := CreateNetNSWithoutSwitching(actorNSName)
|
||||
actorNS, err := netns.CreateNamed(actorNSName)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("while creating the actor netns %s: %w", actorNSName, err)
|
||||
}
|
||||
defer func() {
|
||||
if retErr != nil {
|
||||
actorNS.Close()
|
||||
_ = removeNamedNetNS(actorNSName)
|
||||
_ = netns.RemoveNamed(actorNSName)
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -109,7 +107,7 @@ func SetupSandboxNetwork(ctx context.Context, cfg SandboxNetworkConfig) (_ *Sand
|
||||
defer func() {
|
||||
if retErr != nil {
|
||||
outer.Close()
|
||||
_ = removeNamedNetNS(SandboxGatewayNetNSName(cfg.ActorUID))
|
||||
_ = netns.RemoveNamed(SandboxGatewayNetNSName(cfg.ActorUID))
|
||||
}
|
||||
}()
|
||||
atunnelNS = outer
|
||||
@@ -127,21 +125,21 @@ func SetupSandboxNetwork(ctx context.Context, cfg SandboxNetworkConfig) (_ *Sand
|
||||
|
||||
// setupVethPair creates the gateway namespace and the veth pair joining it to
|
||||
// actorNS. The caller owns the returned handle and its name.
|
||||
func setupVethPair(ctx context.Context, cfg SandboxNetworkConfig, actorNS netns.NsHandle) (_ netns.NsHandle, retErr error) {
|
||||
func setupVethPair(ctx context.Context, cfg SandboxNetworkConfig, actorNS netns.Handle) (_ netns.Handle, retErr error) {
|
||||
gatewayNSName := SandboxGatewayNetNSName(cfg.ActorUID)
|
||||
outer, err := CreateNetNSWithoutSwitching(gatewayNSName)
|
||||
outer, err := netns.CreateNamed(gatewayNSName)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("while creating the outer netns %s: %w", gatewayNSName, err)
|
||||
}
|
||||
defer func() {
|
||||
if retErr != nil {
|
||||
outer.Close()
|
||||
_ = removeNamedNetNS(gatewayNSName)
|
||||
_ = netns.RemoveNamed(gatewayNSName)
|
||||
}
|
||||
}()
|
||||
|
||||
// Keep the kernel-owned peer outside gVisor's namespace.
|
||||
if err := NetNSDo(ctx, outer, func(context.Context) error {
|
||||
if err := netns.Do(ctx, outer, func(context.Context) error {
|
||||
veth := &netlink.Veth{
|
||||
LinkAttrs: netlink.LinkAttrs{Name: gatewayVethName},
|
||||
PeerName: ActorVethName,
|
||||
@@ -170,7 +168,7 @@ func setupVethPair(ctx context.Context, cfg SandboxNetworkConfig, actorNS netns.
|
||||
}
|
||||
|
||||
// gVisor imports these addresses and routes into its network stack.
|
||||
if err := NetNSDo(ctx, actorNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, actorNS, func(context.Context) error {
|
||||
// Loopback lets the actor reach its own address.
|
||||
if err := linkUp("lo"); err != nil {
|
||||
return err
|
||||
@@ -204,11 +202,11 @@ func linkUp(name string) error {
|
||||
}
|
||||
|
||||
// setupGatewaySide brings up lo and puts atunnel in front of the actor's TCP.
|
||||
func setupGatewaySide(ctx context.Context, ns netns.NsHandle, egressPort uint16) error {
|
||||
if err := NetNSDo(ctx, ns, func(context.Context) error {
|
||||
func setupGatewaySide(ctx context.Context, ns netns.Handle, egressPort uint16) error {
|
||||
if err := netns.Do(ctx, ns, func(context.Context) error {
|
||||
// atunnel answers the actor's DNS on 53, and the worker holds no
|
||||
// CAP_NET_BIND_SERVICE.
|
||||
if err := AllowUnprivilegedPorts(); err != nil {
|
||||
if err := netns.AllowUnprivilegedPorts(); err != nil {
|
||||
return err
|
||||
}
|
||||
return linkUp("lo")
|
||||
@@ -221,7 +219,7 @@ func setupGatewaySide(ctx context.Context, ns netns.NsHandle, egressPort uint16)
|
||||
// installEgressRedirect redirects TCP egress to atunnel, excluding the sandbox's
|
||||
// own /30: that keeps ingress replies and DNS over TCP to the gateway off the
|
||||
// redirect, so the relay serves them on its own listener.
|
||||
func installEgressRedirect(ns netns.NsHandle, egressPort uint16) error {
|
||||
func installEgressRedirect(ns netns.Handle, egressPort uint16) error {
|
||||
if egressPort == 0 {
|
||||
return fmt.Errorf("actornet: atunnel egress port is required")
|
||||
}
|
||||
@@ -298,39 +296,13 @@ func CleanupSandboxNetwork(network *SandboxNetwork) error {
|
||||
}
|
||||
// Deleting the namespaces takes any veth pair with them.
|
||||
for _, name := range []string{nodepath.ActorNetNSName(network.ActorUID), SandboxGatewayNetNSName(network.ActorUID)} {
|
||||
if err := removeNamedNetNS(name); err != nil {
|
||||
if err := netns.RemoveNamed(name); err != nil {
|
||||
errs = errors.Join(errs, fmt.Errorf("while deleting netns %s: %w", name, err))
|
||||
}
|
||||
}
|
||||
return errs
|
||||
}
|
||||
|
||||
// ListenInNetNS opens wildcard TCP listeners inside ns.
|
||||
// Sockets retain their namespace and can be served from another namespace.
|
||||
func ListenInNetNS(ctx context.Context, ns netns.NsHandle, ports []uint16) (_ []net.Listener, retErr error) {
|
||||
var listeners []net.Listener
|
||||
defer func() {
|
||||
if retErr != nil {
|
||||
for _, l := range listeners {
|
||||
_ = l.Close()
|
||||
}
|
||||
}
|
||||
}()
|
||||
if err := NetNSDo(ctx, ns, func(context.Context) error {
|
||||
for _, port := range ports {
|
||||
l, err := net.Listen("tcp", fmt.Sprintf("0.0.0.0:%d", port))
|
||||
if err != nil {
|
||||
return fmt.Errorf("while listening on port %d: %w", port, err)
|
||||
}
|
||||
listeners = append(listeners, l)
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return listeners, nil
|
||||
}
|
||||
|
||||
// egressServer serves one actor's captured connections. Satisfied by
|
||||
// atunnel.Egress; an interface so this package does not depend on it.
|
||||
type egressServer interface {
|
||||
@@ -339,8 +311,8 @@ type egressServer interface {
|
||||
|
||||
// ServeSandboxEgress serves redirected TCP in the gateway namespace.
|
||||
// Closing the returned listeners stops accepting new connections.
|
||||
func serveSandboxEgress(ctx context.Context, e egressServer, actorUID string, ns netns.NsHandle, ports []uint16) ([]io.Closer, []func(), error) {
|
||||
listeners, err := ListenInNetNS(ctx, ns, ports)
|
||||
func serveSandboxEgress(ctx context.Context, e egressServer, actorUID string, ns netns.Handle, ports []uint16) ([]io.Closer, []func(), error) {
|
||||
listeners, err := netns.Listen(ctx, ns, ports)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("while opening actor egress listeners: %w", err)
|
||||
}
|
||||
@@ -368,117 +340,6 @@ func serveSandboxEgress(ctx context.Context, e egressServer, actorUID string, ns
|
||||
return closers, serve, nil
|
||||
}
|
||||
|
||||
// closerFunc adapts a cancel function to io.Closer, so a caller takes a
|
||||
// sandbox's sockets and the work behind them down as one list.
|
||||
type closerFunc func() error
|
||||
|
||||
func (f closerFunc) Close() error { return f() }
|
||||
|
||||
// withNetNS switches to targetNS, calls run, then restores the original namespace.
|
||||
// run can call restore to switch back and unlock the OS thread before returning.
|
||||
// Calling restore again after it succeeds has no effect.
|
||||
//
|
||||
// A separate goroutine lets us leave the thread locked if restoration fails.
|
||||
// Go then discards that thread when the goroutine exits.
|
||||
func withNetNS(targetNS netns.NsHandle, run func(restore func() error) error) error {
|
||||
var resultErr error
|
||||
var done sync.WaitGroup
|
||||
done.Add(1)
|
||||
go func() {
|
||||
defer done.Done()
|
||||
runtime.LockOSThread()
|
||||
originalNS, err := netns.Get()
|
||||
if err != nil {
|
||||
runtime.UnlockOSThread()
|
||||
resultErr = fmt.Errorf("while reading the current netns: %w", err)
|
||||
return
|
||||
}
|
||||
defer originalNS.Close()
|
||||
if err := netns.Set(targetNS); err != nil {
|
||||
runtime.UnlockOSThread()
|
||||
resultErr = fmt.Errorf("while entering the actor netns: %w", err)
|
||||
return
|
||||
}
|
||||
|
||||
restored := false
|
||||
restore := func() error {
|
||||
if restored {
|
||||
return nil
|
||||
}
|
||||
if err := netns.Set(originalNS); err != nil {
|
||||
return fmt.Errorf("while restoring the worker netns: %w", err)
|
||||
}
|
||||
runtime.UnlockOSThread()
|
||||
restored = true
|
||||
return nil
|
||||
}
|
||||
|
||||
resultErr = run(restore)
|
||||
if err := restore(); err != nil {
|
||||
resultErr = err
|
||||
}
|
||||
}()
|
||||
done.Wait()
|
||||
return resultErr
|
||||
}
|
||||
|
||||
// NetNSDialer dials TCP or UDP IP literals in ns, pinning a thread only until
|
||||
// the socket is created.
|
||||
func NetNSDialer(ns netns.NsHandle) func(context.Context, string, string) (net.Conn, error) {
|
||||
return func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateNetNSDialTarget(network, addr); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var conn net.Conn
|
||||
dialErr := withNetNS(ns, func(restore func() error) error {
|
||||
// Only creating the socket needs the namespace, and
|
||||
// ControlContext runs once it exists: restore there rather than
|
||||
// holding the thread for the whole connect.
|
||||
socketCreated := false
|
||||
dialer := net.Dialer{ControlContext: func(context.Context, string, string, syscall.RawConn) error {
|
||||
if socketCreated {
|
||||
return errors.New("sandbox dial cannot recreate its socket outside the namespace")
|
||||
}
|
||||
socketCreated = true
|
||||
return restore()
|
||||
}}
|
||||
var err error
|
||||
conn, err = dialer.DialContext(ctx, network, addr)
|
||||
return err
|
||||
})
|
||||
if dialErr != nil || ctx.Err() != nil {
|
||||
if conn != nil {
|
||||
_ = conn.Close()
|
||||
}
|
||||
if dialErr != nil {
|
||||
return nil, dialErr
|
||||
}
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
}
|
||||
|
||||
func validateNetNSDialTarget(network, addr string) error {
|
||||
switch network {
|
||||
case "tcp", "tcp4", "tcp6", "udp", "udp4", "udp6":
|
||||
default:
|
||||
return net.UnknownNetworkError(network)
|
||||
}
|
||||
hostname, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := netip.ParseAddr(hostname); err != nil {
|
||||
return fmt.Errorf("NetNSDialer supports only IP literals (got %q): %w", hostname, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SandboxSession owns a sandbox's network and serving sockets.
|
||||
type SandboxSession struct {
|
||||
Network *SandboxNetwork
|
||||
@@ -492,7 +353,7 @@ type SandboxSession struct {
|
||||
|
||||
// ServeSandbox builds a sandbox's network and serves egress and DNS from its
|
||||
// gateway namespace. A nil server leaves that unserved, which fails closed.
|
||||
func ServeSandbox(ctx context.Context, cfg SandboxNetworkConfig, egress egressServer, dns dnsServer) (_ *SandboxSession, retErr error) {
|
||||
func ServeSandbox(ctx context.Context, cfg SandboxNetworkConfig, egress egressServer, resolver dns.Server) (_ *SandboxSession, retErr error) {
|
||||
network, err := SetupSandboxNetwork(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -505,8 +366,8 @@ func ServeSandbox(ctx context.Context, cfg SandboxNetworkConfig, egress egressSe
|
||||
}()
|
||||
|
||||
var serve []func()
|
||||
if dns != nil {
|
||||
closers, serveDNS, err := serveSandboxDNS(ctx, dns, network.GatewayNetNS, cfg.DNSPort)
|
||||
if resolver != nil {
|
||||
closers, serveDNS, err := dns.Serve(ctx, resolver, network.GatewayNetNS, cfg.DNSPort)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -635,8 +496,8 @@ func (s *SandboxSession) Dialer() func(context.Context, string, string) (net.Con
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("while retaining the sandbox namespace: %w", err)
|
||||
}
|
||||
ns := netns.NsHandle(fd)
|
||||
ns := netns.Handle(fd)
|
||||
defer ns.Close()
|
||||
return NetNSDialer(ns)(ctx, network, address)
|
||||
return netns.Dialer(ns)(ctx, network, address)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,8 +14,8 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
// These integration tests use an external package to avoid an import cycle:
|
||||
// atunnel imports ateomnet.
|
||||
// These integration tests use an external package so that ateomnet's own test
|
||||
// binary does not depend on atunnel.
|
||||
package ateomnet_test
|
||||
|
||||
import (
|
||||
@@ -26,6 +26,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
"github.com/agent-substrate/substrate/internal/atunnel"
|
||||
"github.com/agent-substrate/substrate/internal/roottest"
|
||||
)
|
||||
@@ -44,9 +45,9 @@ func TestSandboxEgressReachesAtunnelOnAnyPort(t *testing.T) {
|
||||
}
|
||||
t.Cleanup(func() { ateomnet.CleanupSandboxNetwork(n) })
|
||||
|
||||
listeners, err := ateomnet.ListenInNetNS(ctx, n.GatewayNetNS, []uint16{egressPort})
|
||||
listeners, err := netns.Listen(ctx, n.GatewayNetNS, []uint16{egressPort})
|
||||
if err != nil {
|
||||
t.Fatalf("ListenInNetNS: %v", err)
|
||||
t.Fatalf("netns.Listen: %v", err)
|
||||
}
|
||||
defer listeners[0].Close()
|
||||
|
||||
@@ -70,7 +71,7 @@ func TestSandboxEgressReachesAtunnelOnAnyPort(t *testing.T) {
|
||||
go accept()
|
||||
|
||||
for _, want := range []string{"93.184.216.34:443", "93.184.216.34:8080", "93.184.216.34:9999"} {
|
||||
if err := ateomnet.NetNSDo(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
c, err := net.Dial("tcp", want)
|
||||
if err != nil {
|
||||
return err
|
||||
|
||||
@@ -32,55 +32,12 @@ import (
|
||||
"github.com/agent-substrate/substrate/internal/nodepath"
|
||||
"github.com/vishvananda/netlink"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
"github.com/agent-substrate/substrate/internal/roottest"
|
||||
"github.com/vishvananda/netns"
|
||||
)
|
||||
|
||||
const testEgressPort = 15001
|
||||
|
||||
func TestValidateNetNSDialTarget(t *testing.T) {
|
||||
for _, target := range []struct {
|
||||
network, address string
|
||||
wantErr bool
|
||||
}{
|
||||
{"tcp", "127.0.0.1:80", false},
|
||||
{"tcp4", "127.0.0.1:80", false},
|
||||
{"tcp6", "[::1]:80", false},
|
||||
{"udp", "127.0.0.1:53", false},
|
||||
{"udp4", "127.0.0.1:53", false},
|
||||
{"udp6", "[fe80::1%eth0]:53", false},
|
||||
{"tcp", "localhost:80", true},
|
||||
{"tcp", ":80", true},
|
||||
{"tcp", "127.0.0.1", true},
|
||||
{"unix", "/tmp/socket", true},
|
||||
{"ip", "127.0.0.1:80", true},
|
||||
} {
|
||||
t.Run(target.network+"/"+target.address, func(t *testing.T) {
|
||||
err := validateNetNSDialTarget(target.network, target.address)
|
||||
if (err != nil) != target.wantErr {
|
||||
t.Errorf("validateNetNSDialTarget = %v, want error: %t", err, target.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
if err := validateNetNSDialTarget("unix", "/tmp/socket"); !errors.Is(err, net.UnknownNetworkError("unix")) {
|
||||
t.Errorf("unsupported network: got %v, want UnknownNetworkError", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNetNSDialerRejectsNonIPTargets(t *testing.T) {
|
||||
for _, target := range []struct{ network, address string }{
|
||||
{"tcp", "localhost:80"},
|
||||
{"tcp", ":80"},
|
||||
{"tcp", "127.0.0.1"},
|
||||
{"unix", "/tmp/socket"},
|
||||
} {
|
||||
if conn, err := NetNSDialer(-1)(context.Background(), target.network, target.address); err == nil {
|
||||
_ = conn.Close()
|
||||
t.Errorf("accepted %s %s", target.network, target.address)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSandboxSessionDialerAfterClose(t *testing.T) {
|
||||
session := &SandboxSession{}
|
||||
dial := session.Dialer()
|
||||
@@ -128,14 +85,6 @@ func TestSandboxSessionDialerConcurrentClose(t *testing.T) {
|
||||
workers.Wait()
|
||||
}
|
||||
|
||||
func TestNetNSDialerCanceledContext(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if _, err := NetNSDialer(-1)(ctx, "tcp", "127.0.0.1:1"); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("canceled dial: got %v, want cancellation", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetupSandboxNetwork(t *testing.T) {
|
||||
roottest.Require(t, "creates network namespaces")
|
||||
ctx := context.Background()
|
||||
@@ -158,7 +107,7 @@ func TestSetupSandboxNetwork(t *testing.T) {
|
||||
|
||||
// The actor's app, bound where a real one binds, inside its namespace.
|
||||
var lis net.Listener
|
||||
if err := NetNSDo(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
l, err := net.Listen("tcp", net.JoinHostPort(ActorVethIP, "80"))
|
||||
lis = l
|
||||
return err
|
||||
@@ -175,7 +124,7 @@ func TestSetupSandboxNetwork(t *testing.T) {
|
||||
|
||||
// Reaching each actor is a matter of which namespace the dial is made from.
|
||||
for uid, a := range actors {
|
||||
client := &http.Client{Transport: &http.Transport{DialContext: NetNSDialer(a.net.RuntimeNetNS)}, Timeout: 5 * time.Second}
|
||||
client := &http.Client{Transport: &http.Transport{DialContext: netns.Dialer(a.net.RuntimeNetNS)}, Timeout: 5 * time.Second}
|
||||
resp, err := client.Get((&url.URL{Scheme: "http", Host: net.JoinHostPort(ActorVethIP, "80")}).String())
|
||||
if err != nil {
|
||||
t.Fatalf("reaching actor %s: %v", uid, err)
|
||||
@@ -207,7 +156,7 @@ func TestActorEgressIsFailClosedWithoutAtunnel(t *testing.T) {
|
||||
t.Cleanup(func() { CleanupSandboxNetwork(n) })
|
||||
|
||||
for _, destination := range []string{"93.184.216.34:443", "93.184.216.34:8080"} {
|
||||
if err := NetNSDo(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
c, err := net.DialTimeout("tcp", destination, 3*time.Second)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -231,7 +180,7 @@ func TestIngressCrossesThePairWhileEgressIsCaptured(t *testing.T) {
|
||||
t.Cleanup(func() { CleanupSandboxNetwork(n) })
|
||||
|
||||
var app net.Listener
|
||||
if err := NetNSDo(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
l, e := net.Listen("tcp", net.JoinHostPort(ActorVethIP, "80"))
|
||||
app = l
|
||||
return e
|
||||
@@ -250,7 +199,7 @@ func TestIngressCrossesThePairWhileEgressIsCaptured(t *testing.T) {
|
||||
}
|
||||
}()
|
||||
|
||||
dial := NetNSDialer(n.GatewayNetNS)
|
||||
dial := netns.Dialer(n.GatewayNetNS)
|
||||
c, err := dial(ctx, "tcp", net.JoinHostPort(ActorVethIP, "80"))
|
||||
if err != nil {
|
||||
t.Fatalf("ingress dial: %v", err)
|
||||
@@ -289,7 +238,7 @@ func TestSetupSandboxNetworkWithoutVeth(t *testing.T) {
|
||||
}
|
||||
|
||||
// No veth was built, so nothing but lo is here until the tap arrives.
|
||||
if err := NetNSDo(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, n.RuntimeNetNS, func(context.Context) error {
|
||||
links, err := netlink.LinkList()
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -332,7 +281,7 @@ func TestSetupSucceedsOverALeftoverNamespace(t *testing.T) {
|
||||
}
|
||||
t.Cleanup(func() { CleanupSandboxNetwork(second) })
|
||||
|
||||
if err := NetNSDo(ctx, second.RuntimeNetNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, second.RuntimeNetNS, func(context.Context) error {
|
||||
if _, err := netlink.LinkByName(ActorVethName); err != nil {
|
||||
return fmt.Errorf("actor interface missing after reuse: %w", err)
|
||||
}
|
||||
@@ -356,7 +305,7 @@ func TestActorUDPHasNowhereToGoBeyondTheNamespacePair(t *testing.T) {
|
||||
t.Cleanup(func() { CleanupSandboxNetwork(n) })
|
||||
|
||||
for _, destination := range []string{"93.184.216.34:443", "93.184.216.34:53"} {
|
||||
if err := NetNSDo(ctx, n.GatewayNetNS, func(context.Context) error {
|
||||
if err := netns.Do(ctx, n.GatewayNetNS, func(context.Context) error {
|
||||
c, err := net.Dial("udp", destination)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -389,56 +338,6 @@ func TestCleanupClosesEachDescriptorOnce(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// stoppableDNS records that its serving contexts were canceled.
|
||||
type stoppableDNS struct{ packet, stream chan struct{} }
|
||||
|
||||
func (d *stoppableDNS) ServePacket(ctx context.Context, pc net.PacketConn) error {
|
||||
<-ctx.Done()
|
||||
close(d.packet)
|
||||
return pc.Close()
|
||||
}
|
||||
|
||||
func (d *stoppableDNS) Serve(ctx context.Context, l net.Listener) error {
|
||||
<-ctx.Done()
|
||||
close(d.stream)
|
||||
return l.Close()
|
||||
}
|
||||
|
||||
func TestClosingSandboxDNSStopsServing(t *testing.T) {
|
||||
roottest.Require(t, "creates network namespaces")
|
||||
network, err := SetupSandboxNetwork(context.Background(), SandboxNetworkConfig{
|
||||
ActorUID: "dns-teardown",
|
||||
EgressPort: 15001,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = CleanupSandboxNetwork(network) }()
|
||||
|
||||
relay := &stoppableDNS{packet: make(chan struct{}), stream: make(chan struct{})}
|
||||
closers, serve, err := serveSandboxDNS(context.Background(), relay, network.GatewayNetNS, 53)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, fn := range serve {
|
||||
go fn()
|
||||
}
|
||||
for _, c := range closers {
|
||||
_ = c.Close()
|
||||
}
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
stopped chan struct{}
|
||||
}{{"UDP", relay.packet}, {"TCP", relay.stream}} {
|
||||
select {
|
||||
case <-tc.stopped:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Errorf("%s serving outlived the sandbox's sockets", tc.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// slowDNS holds its serving goroutines open until released, so a test can tell
|
||||
// whether Close waits for them or merely closes their sockets.
|
||||
type slowDNS struct{ release chan struct{} }
|
||||
|
||||
@@ -31,10 +31,9 @@ import (
|
||||
"github.com/google/nftables/binaryutil"
|
||||
"github.com/google/nftables/expr"
|
||||
"github.com/vishvananda/netlink"
|
||||
"github.com/vishvananda/netns"
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet"
|
||||
"github.com/agent-substrate/substrate/internal/ateomnet/netns"
|
||||
"github.com/agent-substrate/substrate/internal/roottest"
|
||||
)
|
||||
|
||||
@@ -66,7 +65,7 @@ func TestTCPOriginalDestinationPreservesErrno(t *testing.T) {
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
ns := newTestNetNS(t)
|
||||
if err := ateomnet.NetNSDo(context.Background(), ns, func(context.Context) error {
|
||||
if err := netns.Do(context.Background(), ns, func(context.Context) error {
|
||||
loopback, err := netlink.LinkByName("lo")
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -140,7 +139,7 @@ func TestTCPOriginalDestination(t *testing.T) {
|
||||
// From the actor's perspective this is an ordinary connection to
|
||||
// hostIP:targetPort. The worker's PREROUTING rule redirects it before
|
||||
// it reaches the host network stack's local delivery path.
|
||||
clientDone <- ateomnet.NetNSDo(context.Background(), actorNS, func(context.Context) error {
|
||||
clientDone <- netns.Do(context.Background(), actorNS, func(context.Context) error {
|
||||
conn, err := net.DialTimeout("tcp4", net.JoinHostPort(hostIP.String(), fmt.Sprint(targetPort)), 10*time.Second)
|
||||
if err == nil {
|
||||
_ = conn.Close()
|
||||
@@ -192,7 +191,7 @@ func TestTCPOriginalDestinationIPv6(t *testing.T) {
|
||||
|
||||
clientDone := make(chan error, 1)
|
||||
go func() {
|
||||
clientDone <- ateomnet.NetNSDo(context.Background(), actorNS, func(context.Context) error {
|
||||
clientDone <- netns.Do(context.Background(), actorNS, func(context.Context) error {
|
||||
conn, err := net.DialTimeout("tcp6", net.JoinHostPort(hostIP.String(), fmt.Sprint(targetPort)), 10*time.Second)
|
||||
if err == nil {
|
||||
_ = conn.Close()
|
||||
@@ -231,7 +230,7 @@ func TestTCPOriginalDestinationIPv6(t *testing.T) {
|
||||
func withTestWorkerNS(t *testing.T, fn func()) {
|
||||
t.Helper()
|
||||
workerNS := newTestNetNS(t)
|
||||
if err := ateomnet.NetNSDo(context.Background(), workerNS, func(context.Context) error {
|
||||
if err := netns.Do(context.Background(), workerNS, func(context.Context) error {
|
||||
fn()
|
||||
return nil
|
||||
}); err != nil {
|
||||
@@ -239,10 +238,10 @@ func withTestWorkerNS(t *testing.T, fn func()) {
|
||||
}
|
||||
}
|
||||
|
||||
func newTestNetNS(t *testing.T) netns.NsHandle {
|
||||
func newTestNetNS(t *testing.T) netns.Handle {
|
||||
t.Helper()
|
||||
name := fmt.Sprintf("atunnel-original-dst-%d-%d", os.Getpid(), atomic.AddUint64(&testNetNSSequence, 1))
|
||||
ns, err := ateomnet.CreateNetNSWithoutSwitching(name)
|
||||
ns, err := netns.CreateNamed(name)
|
||||
if err != nil {
|
||||
if errors.Is(err, unix.EPERM) || strings.Contains(err.Error(), "operation not permitted") {
|
||||
t.Skipf("needs CAP_SYS_ADMIN to create network namespace: %v", err)
|
||||
@@ -251,14 +250,14 @@ func newTestNetNS(t *testing.T) netns.NsHandle {
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = ns.Close()
|
||||
if err := netns.DeleteNamed(name); err != nil {
|
||||
if err := netns.RemoveNamed(name); err != nil {
|
||||
t.Errorf("deleting test network namespace: %v", err)
|
||||
}
|
||||
})
|
||||
return ns
|
||||
}
|
||||
|
||||
func setupTestVeth(t *testing.T, actorNS netns.NsHandle) (actorIP, hostIP net.IP) {
|
||||
func setupTestVeth(t *testing.T, actorNS netns.Handle) (actorIP, hostIP net.IP) {
|
||||
t.Helper()
|
||||
hostName := fmt.Sprintf("atod%d", os.Getpid())
|
||||
peerName := fmt.Sprintf("atop%d", os.Getpid())
|
||||
@@ -300,7 +299,7 @@ func setupTestVeth(t *testing.T, actorNS netns.NsHandle) (actorIP, hostIP net.IP
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Complete the actor end of the point-to-point link inside its own netns.
|
||||
if err := ateomnet.NetNSDo(context.Background(), actorNS, func(context.Context) error {
|
||||
if err := netns.Do(context.Background(), actorNS, func(context.Context) error {
|
||||
lo, err := netlink.LinkByName("lo")
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -331,7 +330,7 @@ func listenTCP(t *testing.T, hostIP net.IP) net.Listener {
|
||||
return listener
|
||||
}
|
||||
|
||||
func setupTestIPv6Veth(t *testing.T, actorNS netns.NsHandle) (actorIP, hostIP net.IP) {
|
||||
func setupTestIPv6Veth(t *testing.T, actorNS netns.Handle) (actorIP, hostIP net.IP) {
|
||||
t.Helper()
|
||||
hostName := fmt.Sprintf("atod6%d", os.Getpid())
|
||||
peerName := fmt.Sprintf("atop6%d", os.Getpid())
|
||||
@@ -371,7 +370,7 @@ func setupTestIPv6Veth(t *testing.T, actorNS netns.NsHandle) (actorIP, hostIP ne
|
||||
if err := netlink.LinkSetNsFd(peer, int(actorNS)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := ateomnet.NetNSDo(context.Background(), actorNS, func(context.Context) error {
|
||||
if err := netns.Do(context.Background(), actorNS, func(context.Context) error {
|
||||
lo, err := netlink.LinkByName("lo")
|
||||
if err != nil {
|
||||
return err
|
||||
|
||||
Reference in New Issue
Block a user