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:
Bowei Du
2026-09-28 17:36:52 +00:00
committed by GitHub
parent 6621f2b4b4
commit 22efea18a9
20 changed files with 652 additions and 482 deletions
+2 -1
View File
@@ -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
+2 -2
View File
@@ -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
+5 -5
View File
@@ -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 {
+5 -5
View File
@@ -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
+2 -1
View File
@@ -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 {
@@ -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() }
+78
View File
@@ -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)
}
}
}
-124
View File
@@ -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
}
+135
View File
@@ -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)
}
}
+155
View File
@@ -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")
}
}
+60
View File
@@ -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
View File
@@ -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
+9 -110
View File
@@ -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{} }
+12 -13
View File
@@ -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