diff --git a/apps/daemon/cmd/oac-process-shim/main.go b/apps/daemon/cmd/oac-process-shim/main.go new file mode 100644 index 00000000..c0b1b9e5 --- /dev/null +++ b/apps/daemon/cmd/oac-process-shim/main.go @@ -0,0 +1,17 @@ +// Command oac-process-shim runs a declared sandbox executable from a Session +// view through the process broker, and is the Session's process relay. See +// apps/daemon/internal/processshim. +package main + +import ( + "os" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" +) + +func main() { + if processshim.Relaying() { + os.Exit(processshim.Relay()) + } + os.Exit(processshim.Run(processshim.SocketPath)) +} diff --git a/apps/daemon/internal/agent/harness.go b/apps/daemon/internal/agent/harness.go index 2528f534..050fa4b4 100644 --- a/apps/daemon/internal/agent/harness.go +++ b/apps/daemon/internal/agent/harness.go @@ -113,8 +113,11 @@ const ( ViewShimName = "bin" // ViewHomeName is the Session home under ViewPrivateRoot. ViewHomeName = "home" - // ViewRunName is the process broker's socket directory under ViewPrivateRoot. + // ViewRunName is the process relay's socket directory under ViewPrivateRoot. ViewRunName = "run" + // ViewRelayName is the process relay's name in the shim directory, which + // no shim takes. + ViewRelayName = "oac-process-shim" // ViewProcRoot and ViewDevRoot are the view's own /proc and minimal /dev. ViewProcRoot = "/proc" ViewDevRoot = "/dev" @@ -335,7 +338,7 @@ func (v View) Validate() error { } } for i, n := range v.Shims { - if !isPathComponent(n) || slices.Contains(v.Shims[:i], n) { + if !isPathComponent(n) || n == ViewRelayName || slices.Contains(v.Shims[:i], n) { return invalidView("shim %q", n) } } diff --git a/apps/daemon/internal/agent/view_test.go b/apps/daemon/internal/agent/view_test.go index 45384297..b656806e 100644 --- a/apps/daemon/internal/agent/view_test.go +++ b/apps/daemon/internal/agent/view_test.go @@ -32,6 +32,7 @@ func TestViewValidate(t *testing.T) { "mask in /proc": func(v *agent.View) { v.Masks[0].Path = "/proc/cpuinfo" }, "unclean view path": func(v *agent.View) { v.Masks[0].Path = "/etc/../etc/harness" }, "duplicate shim": func(v *agent.View) { v.Shims = append(v.Shims, "git") }, + "shim named as the relay": func(v *agent.View) { v.Shims = append(v.Shims, agent.ViewRelayName) }, "forwarded assignment": func(v *agent.View) { v.ForwardEnv = append(v.ForwardEnv, "A=B") }, "forwarded broker variable": func(v *agent.View) { v.ForwardEnv = append(v.ForwardEnv, "PATH") }, "forwarded proxy variable": func(v *agent.View) { v.ForwardEnv = append(v.ForwardEnv, "https_proxy") }, diff --git a/apps/daemon/internal/agenthost/broker.go b/apps/daemon/internal/agenthost/broker.go index de5d0866..1df6d799 100644 --- a/apps/daemon/internal/agenthost/broker.go +++ b/apps/daemon/internal/agenthost/broker.go @@ -11,7 +11,7 @@ import ( // over the Process service. One broker serves a Session from its first launch // until teardown. type processBroker interface { - // Start begins serving the Session's run directory. + // Start begins serving the Session. Start(brokerConfig) error // Close cancels and releases the remote operations that remain and stops // serving. @@ -20,9 +20,6 @@ type processBroker interface { // brokerConfig is what a Session's broker serves. type brokerConfig struct { - // RunDir is the host directory the view presents read-only at - // agent.ViewPrivateRoot/agent.ViewRunName. - RunDir string // UID and GID are the Session's. UID, GID uint32 // Names maps each shim name to the program it runs in the sandbox, found diff --git a/apps/daemon/internal/agenthost/doc.go b/apps/daemon/internal/agenthost/doc.go index 26014d6b..1c6e7e89 100644 --- a/apps/daemon/internal/agenthost/doc.go +++ b/apps/daemon/internal/agenthost/doc.go @@ -17,10 +17,10 @@ // the gateway listening in the view's network namespace. // // Each view presents the closure directories read-only and executable, the -// Session home read-write and noexec, the broker's run directory read-only, -// the agent host's /etc/passwd, group, hosts, resolv.conf and nsswitch.conf, -// the agent host's CA directory at its host path, then the adapter's overlays -// and masks and the process shim. Everything else is the world. +// Session home read-write and noexec, the agent host's /etc/passwd, group, +// hosts, resolv.conf and nsswitch.conf, the agent host's CA directory at its +// host path, then the adapter's overlays and masks and the process shim with +// its relay. Everything else is the world. // // The agent host owns the Session's Link attachment: it opens each stream // with the Session's binding, renews the lease and fails the Session when the diff --git a/apps/daemon/internal/agenthost/launch_linux.go b/apps/daemon/internal/agenthost/launch_linux.go index a87ebb1c..b38e63dd 100644 --- a/apps/daemon/internal/agenthost/launch_linux.go +++ b/apps/daemon/internal/agenthost/launch_linux.go @@ -243,7 +243,7 @@ func (s *session) startBroker(grace time.Duration) error { paths[p] = p } b := s.deps.broker() - err := b.Start(brokerConfig{RunDir: s.dir.entry(runEntry), UID: s.uid, GID: s.uid, Names: names, Paths: paths, + err := b.Start(brokerConfig{UID: s.uid, GID: s.uid, Names: names, Paths: paths, Pass: slices.Clone(view.ForwardEnv), Sandbox: s.in.Environment.Sandbox, Tool: s.in.Environment.Tool, Dial: s.openProcess, CancelGrace: grace}) if err != nil { @@ -255,7 +255,7 @@ func (s *session) startBroker(grace time.Duration) error { return nil } -// spec builds the view: the closure, home and run directories, the agent +// spec builds the view: the closure and home directories, the agent // host's /etc files and CA directory, the adapter's overlays and masks, the // shim and the gateway in the view's network namespace. func (s *session) spec(viewCtx context.Context, world *worldfs.World, opts clirunner.StartOptions, ends *stdio) sessionview.Spec { @@ -264,9 +264,7 @@ func (s *session) spec(viewCtx context.Context, world *worldfs.World, opts cliru for _, m := range view.Closure { private = append(private, sessionview.PrivateDir{Name: m.Name, HostDir: m.HostDir, Exec: true}) } - private = append(private, - sessionview.PrivateDir{Name: agent.ViewHomeName, HostDir: s.dir.entry(homeEntry), Writable: true}, - sessionview.PrivateDir{Name: agent.ViewRunName, HostDir: s.dir.entry(runEntry)}) + private = append(private, sessionview.PrivateDir{Name: agent.ViewHomeName, HostDir: s.dir.entry(homeEntry), Writable: true}) var overlays []sessionview.Overlay for _, name := range etcFiles { overlays = append(overlays, sessionview.Overlay{Path: "/etc/" + name, Source: s.dir.entry(etcEntry, name)}) diff --git a/apps/daemon/internal/agenthost/sessiondir_linux.go b/apps/daemon/internal/agenthost/sessiondir_linux.go index 42d1ae8c..1602e63f 100644 --- a/apps/daemon/internal/agenthost/sessiondir_linux.go +++ b/apps/daemon/internal/agenthost/sessiondir_linux.go @@ -47,7 +47,6 @@ func freeUID(id uint32) { // The entries of a Session directory. const ( homeEntry = agent.ViewHomeName // the Session home, owned by the Session uid - runEntry = agent.ViewRunName // the process broker's run directory etcEntry = "etc" // the /etc files maskEntry = "mask" // an empty file and an empty directory that masks present stagingEntry = "staging" // sessionview's staging parent @@ -62,8 +61,8 @@ func (d sessionDir) entry(name ...string) string { return filepath.Join(append([]string{string(d)}, name...)...) } -// createSessionDir creates the Session directory with its home, run, etc, -// mask and staging entries. +// createSessionDir creates the Session directory with its home, etc, mask +// and staging entries. func createSessionDir(stateDir string, id sandboxwire.ID, uid uint32) (sessionDir, error) { parent := sessionsDir(stateDir) if err := os.MkdirAll(parent, 0o700); err != nil { @@ -86,7 +85,7 @@ func (d sessionDir) populate(uid uint32) error { dirs := []struct { name string mode fs.FileMode - }{{homeEntry, 0o700}, {runEntry, 0o755}, {etcEntry, 0o755}, {maskEntry, 0o755}, {filepath.Join(maskEntry, "dir"), 0o555}, {stagingEntry, 0o700}} + }{{homeEntry, 0o700}, {etcEntry, 0o755}, {maskEntry, 0o755}, {filepath.Join(maskEntry, "dir"), 0o555}, {stagingEntry, 0o700}} for _, e := range dirs { if err := os.Mkdir(d.entry(e.name), e.mode); err != nil { return err diff --git a/apps/daemon/internal/processbroker/broker_linux.go b/apps/daemon/internal/processbroker/broker_linux.go new file mode 100644 index 00000000..332c49fd --- /dev/null +++ b/apps/daemon/internal/processbroker/broker_linux.go @@ -0,0 +1,169 @@ +//go:build linux + +package processbroker + +import ( + "bufio" + "context" + "errors" + "fmt" + "log/slog" + "net" + "sync" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// Broker serves one Session's shim invocations, which its relay hands it. +type Broker struct { + cfg Config + log *slog.Logger + conn *net.UnixConn + link *link + + ctx context.Context + cancel context.CancelFunc + wg sync.WaitGroup // the reader and the invocations + sendMu sync.Mutex + done chan struct{} + + mu sync.Mutex + invs map[uint64]*invocation + lastID uint64 + closing bool + err error + + closeOnce sync.Once +} + +// Start serves the relay at cfg.Relay until Close or the relay's loss. +func Start(cfg Config) (*Broker, error) { + if err := cfg.validate(); err != nil { + return nil, err + } + log := cfg.Logger + if log == nil { + log = slog.Default() + } + c, err := net.FileConn(cfg.Relay) + if err != nil { + return nil, fmt.Errorf("processbroker: relay connection: %w", err) + } + uc, ok := c.(*net.UnixConn) + if !ok { + c.Close() + return nil, fmt.Errorf("%w: the relay connection is not a Unix socket", ErrInvalidConfig) + } + ctx, cancel := context.WithCancel(context.Background()) + b := &Broker{ + cfg: cfg, log: log, conn: uc, link: newLink(cfg.Dial, log), + ctx: ctx, cancel: cancel, done: make(chan struct{}), invs: map[uint64]*invocation{}, + } + b.wg.Add(1) + go b.read() + return b, nil +} + +// Close ends every invocation and the relay connection, and waits for the +// invocations. It never waits on the relay: the relay answers a waiting shim +// with 255. Remote operations are left to the Session. +func (b *Broker) Close() error { + b.closeOnce.Do(func() { + b.mu.Lock() + b.closing = true + b.mu.Unlock() + b.cancel() + b.conn.CloseWrite() + b.conn.Close() + b.wg.Wait() + b.link.close() + }) + return nil +} + +// Done closes when the broker stops serving: after Close, or when the relay +// is lost. +func (b *Broker) Done() <-chan struct{} { return b.done } + +// Err returns an error wrapping ErrRelayLost once the relay was lost before +// Close, and nil otherwise. +func (b *Broker) Err() error { + b.mu.Lock() + defer b.mu.Unlock() + return b.err +} + +// read dispatches the relay's messages until the connection ends or breaks +// the IPC. It reads without a control buffer, so the kernel closes any +// descriptor the relay attaches, and it never blocks on an invocation. +func (b *Broker) read() { + defer b.wg.Done() + r := bufio.NewReaderSize(b.conn, 64<<10) + var err error + for err == nil { + var f sandboxwire.Frame + if f, err = sandboxwire.ReadFrame(r, processshim.MaxFrameBytes); err != nil { + break + } + var m processshim.RelayMessage + if m, err = processshim.DecodeRelay(f); err == nil { + err = b.dispatch(m) + } + } + b.mu.Lock() + lost := !b.closing + if lost { + b.err = fmt.Errorf("%w: %w", ErrRelayLost, err) + } + b.mu.Unlock() + if lost { + b.log.Error("process relay lost", "error", err) + } + b.cancel() + b.conn.Close() + close(b.done) +} + +func (b *Broker) dispatch(m processshim.RelayMessage) error { + b.mu.Lock() + defer b.mu.Unlock() + id := m.Invocation() + if open, ok := m.(processshim.Open); ok { + switch { + case id <= b.lastID: + return fmt.Errorf("%w: open of invocation %d after %d", processshim.ErrProtocol, id, b.lastID) + case len(b.invs) >= processshim.MaxInvocations: + return fmt.Errorf("%w: more than %d invocations", processshim.ErrProtocol, processshim.MaxInvocations) + } + b.lastID = id + inv := b.newInvocation(open) + b.invs[id] = inv + b.wg.Add(1) + go inv.serve() + return nil + } + inv := b.invs[id] + switch { + case inv != nil: + return inv.receive(m) + case id > b.lastID: + return fmt.Errorf("%w: message for unopened invocation %d", processshim.ErrProtocol, id) + } + return nil // an invocation that has ended +} + +// unregister stops dispatching to the invocation. +func (b *Broker) unregister(id uint64) { + b.mu.Lock() + defer b.mu.Unlock() + delete(b.invs, id) +} + +var errEnded = errors.New("invocation ended") + +func (b *Broker) send(m processshim.BrokerMessage) error { + b.sendMu.Lock() + defer b.sendMu.Unlock() + return sandboxwire.WriteFrame(b.conn, processshim.Frame(m)) +} diff --git a/apps/daemon/internal/processbroker/broker_linux_test.go b/apps/daemon/internal/processbroker/broker_linux_test.go new file mode 100644 index 00000000..bb73c314 --- /dev/null +++ b/apps/daemon/internal/processbroker/broker_linux_test.go @@ -0,0 +1,1299 @@ +//go:build linux + +package processbroker + +import ( + "bufio" + "bytes" + "context" + "errors" + "fmt" + "io" + "net" + "os" + "os/exec" + "path/filepath" + "runtime" + "strconv" + "strings" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/creack/pty" + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/sessionview" + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// The sandbox side is the Linux process service in its own process, +// processserve, because the service reaps every child of its process and +// this binary waits on the shims it starts. Locally the tests build +// processserve with the go command. As root, the relay runs as viewID, as +// in a view. The view tests need root in a privileged container, with +// static binaries: +// +// CGO_ENABLED=0 go test -c -o /tmp/processbroker.test ./apps/daemon/internal/processbroker +// CGO_ENABLED=0 go build -o /tmp/processserve ./apps/sandboxio/testdata/processserve +// CGO_ENABLED=0 go build -o /tmp/oac-process-shim ./apps/daemon/cmd/oac-process-shim +// docker run --rm --cap-add SYS_ADMIN --cap-add NET_ADMIN --device /dev/fuse --security-opt apparmor=unconfined \ +// -e OAC_TEST_SESSIONVIEW=1 -e OAC_TEST_PROCESS_SERVICE=/svc -e OAC_TEST_PROCESS_SHIM=/shim \ +// -v /tmp/processbroker.test:/t.test:ro -v /tmp/processserve:/svc:ro -v /tmp/oac-process-shim:/shim:ro \ +// debian:bookworm-slim /t.test -test.v +const ( + shimSocketEnv = "OAC_TEST_SHIM_SOCKET" + relayEnv = "OAC_TEST_RELAY" + serviceEnv = "OAC_TEST_PROCESS_SERVICE" + shimEnv = "OAC_TEST_PROCESS_SHIM" + viewGateEnv = "OAC_TEST_SESSIONVIEW" + viewID = 1000 +) + +// The test binary is also the relay, the shim through links named after the +// commands, and a Harness in the view. +func TestMain(m *testing.M) { + sessionview.Init() + if os.Getenv(relayEnv) == "1" { + os.Exit(processshim.Relay()) + } + if sock := os.Getenv(shimSocketEnv); sock != "" { + os.Exit(processshim.Run(sock)) + } + if os.Getenv(harnessEnv) == "1" { + os.Exit(runHarness()) + } + code := m.Run() + service.stop() + os.Exit(code) +} + +func TestStreamsAndExitCode(t *testing.T) { + f := newFixture(t, nil) + cmd := f.command("bash", "-c", "echo out; echo err >&2; exit 3") + var stdout, stderr bytes.Buffer + cmd.Stdout, cmd.Stderr = &stdout, &stderr + err := cmd.Run() + if exitCode(err) != 3 || stdout.String() != "out\n" || stderr.String() != "err\n" { + t.Fatalf("Run = %v; stdout %q, stderr %q", err, stdout.String(), stderr.String()) + } +} + +// The shim exits with the leader, and output a background job writes later +// still reaches the shim's stdout. +func TestBackgroundOutputAfterShimExits(t *testing.T) { + f := newFixture(t, nil) + r, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + defer r.Close() + cmd := f.command("sh", "-c", "(sleep 1; echo late) & echo hi") + cmd.Stdout, cmd.Stderr = w, w + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + w.Close() + read := make(chan []byte) + go func() { + data, _ := io.ReadAll(r) + read <- data + }() + if err := cmd.Wait(); err != nil { + t.Fatalf("Wait = %v", err) + } + select { + case data := <-read: + t.Fatalf("pipe closed before the shim exited; read %q", data) + default: + } + select { + case data := <-read: + if string(data) != "hi\nlate\n" { + t.Fatalf("read %q", data) + } + case <-time.After(10 * time.Second): + t.Fatal("pipe still open 10s after the job should have ended") + } +} + +func TestRemoteSignalDeath(t *testing.T) { + f := newFixture(t, nil) + err := f.command("sh", "-c", "kill -TERM $$").Run() + var ee *exec.ExitError + if !errors.As(err, &ee) { + t.Fatalf("Run = %v", err) + } + if ws := ee.Sys().(syscall.WaitStatus); !ws.Signaled() || ws.Signal() != syscall.SIGTERM { + t.Fatalf("status = %v", ws) + } +} + +func TestInterruptReachesRemoteGroup(t *testing.T) { + f := newFixture(t, nil) + cmd := f.command("bash", "-c", "trap 'echo interrupted; exit 7' INT; echo ready; while :; do sleep 0.1; done") + cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} + out := startLine(t, cmd, "ready") + if err := syscall.Kill(-cmd.Process.Pid, syscall.SIGINT); err != nil { + t.Fatal(err) + } + rest, _ := io.ReadAll(out) + if err := cmd.Wait(); exitCode(err) != 7 || string(rest) != "interrupted\n" { + t.Fatalf("Wait = %v; output %q", err, rest) + } +} + +func TestKilledShimCancelsRemoteScope(t *testing.T) { + f := newFixture(t, nil) + cmd := f.command("bash", "-c", "sleep 1000 & echo $!; wait") + out := bufio.NewReader(start(t, cmd)) + line, err := out.ReadString('\n') + if err != nil { + t.Fatal(err) + } + pid := strings.TrimSpace(line) + if _, err := os.Stat("/proc/" + pid); err != nil { + t.Fatalf("remote job %s: %v", pid, err) + } + cmd.Process.Kill() + cmd.Wait() + deadline := time.Now().Add(10 * time.Second) + for { + if _, err := os.Stat("/proc/" + pid); errors.Is(err, os.ErrNotExist) { + return + } + if time.Now().After(deadline) { + t.Fatalf("remote job %s still running 10s after the shim died", pid) + } + time.Sleep(20 * time.Millisecond) + } +} + +func TestTerminalSizeAndMode(t *testing.T) { + f := newFixture(t, nil) + ptm, pts, err := pty.Open() + if err != nil { + t.Fatal(err) + } + defer ptm.Close() + defer pts.Close() + if err := pty.Setsize(ptm, &pty.Winsize{Rows: 30, Cols: 100}); err != nil { + t.Fatal(err) + } + before, err := unix.IoctlGetTermios(int(pts.Fd()), unix.TCGETS) + if err != nil { + t.Fatal(err) + } + cmd := f.command("bash", "-c", `stty size; while [ "$(stty size)" = "30 100" ]; do sleep 0.1; done; stty size`) + cmd.Stdin, cmd.Stdout, cmd.Stderr = pts, pts, pts + cmd.SysProcAttr = &syscall.SysProcAttr{Setsid: true, Setctty: true} + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + var ( + mu sync.Mutex + out bytes.Buffer + ) + done := make(chan struct{}) + go func() { + defer close(done) + buf := make([]byte, 4096) + for { + n, err := ptm.Read(buf) + mu.Lock() + out.Write(buf[:n]) + mu.Unlock() + if err != nil { + return + } + } + }() + output := func() string { + mu.Lock() + defer mu.Unlock() + return out.String() + } + for deadline := time.Now().Add(10 * time.Second); !strings.Contains(output(), "30 100\r\n"); time.Sleep(10 * time.Millisecond) { + if time.Now().After(deadline) { + t.Fatalf("no initial size; output %q", output()) + } + } + // The kernel sends SIGWINCH to the shim, the terminal's foreground group. + if err := pty.Setsize(ptm, &pty.Winsize{Rows: 40, Cols: 120}); err != nil { + t.Fatal(err) + } + if err := cmd.Wait(); err != nil { + t.Fatalf("Wait = %v; output %q", err, output()) + } + after, err := unix.IoctlGetTermios(int(pts.Fd()), unix.TCGETS) + if err != nil { + t.Fatal(err) + } + if *after != *before { + t.Fatalf("terminal mode not restored:\nbefore %+v\nafter %+v", *before, *after) + } + pts.Close() + select { + case <-done: + case <-time.After(10 * time.Second): + t.Fatalf("terminal still open 10s after the exit; output %q", output()) + } + if got := output(); got != "30 100\r\n40 120\r\n" { + t.Fatalf("output %q", got) + } +} + +func TestEnvironmentIsDeclaredOnly(t *testing.T) { + f := newFixture(t, nil) + cmd := f.command("env") + cmd.Env = append(cmd.Env, "KEEP=kept", "LEAK=/x/.oac/run", "OTHER=dropped", "HOME=/local/home") + out, err := cmd.Output() + if err != nil { + t.Fatalf("Output = %v", err) + } + want := "HOME=/home/sandbox\nKEEP=kept\nLANG=C.UTF-8\nPATH=/usr/bin:/bin\nTOOL=1\n" + if string(out) != want || bytes.Contains(out, []byte("/.oac")) { + t.Fatalf("remote environment %q, want %q", out, want) + } +} + +func TestLinkLossKeepsOutputOrdered(t *testing.T) { + f := newFixture(t, first(func(c net.Conn) io.ReadWriteCloser { return &cutConn{Conn: c, left: 64 << 10} })) + out, err := f.command("seq", "1", "2000000").Output() + if err != nil { + t.Fatalf("Output = %v", err) + } + var want []byte + for i := 1; i <= 2000000; i++ { + want = strconv.AppendInt(want, int64(i), 10) + want = append(want, '\n') + } + if !bytes.Equal(out, want) { + t.Fatalf("output differs: %d bytes, want %d", len(out), len(want)) + } + if n := f.dials.Load(); n < 2 { + t.Fatalf("%d dials; the link was never cut", n) + } +} + +// A shim request that never completes holds neither another invocation nor +// Close, and the descriptors it carried close with the relay. +func TestIncompleteRequestHoldsNothing(t *testing.T) { + f := newFixture(t, nil) + c, err := net.DialUnix("unix", nil, &net.UnixAddr{Name: filepath.Join(f.dir, processshim.SocketName), Net: "unix"}) + if err != nil { + t.Fatal(err) + } + defer c.Close() + var p [2]int + if err := unix.Pipe2(p[:], unix.O_CLOEXEC); err != nil { + t.Fatal(err) + } + defer unix.Close(p[0]) + var req bytes.Buffer + if err := sandboxwire.WriteFrame(&req, processshim.Frame(processshim.Request{Version: processshim.Version})); err != nil { + t.Fatal(err) + } + _, _, err = c.WriteMsgUnix(req.Bytes()[:req.Len()/2], unix.UnixRights(p[1], p[1], p[1]), nil) + unix.Close(p[1]) + if err != nil { + t.Fatal(err) + } + if out, err := f.command("sh", "-c", "echo ok").Output(); err != nil || string(out) != "ok\n" { + t.Fatalf("Output = %q, %v", out, err) + } + closed := make(chan struct{}) + go func() { + f.broker.Close() + close(closed) + }() + select { + case <-closed: + case <-time.After(5 * time.Second): + t.Fatal("Close waits for the incomplete request") + } + pfd := []unix.PollFd{{Fd: int32(p[0]), Events: unix.POLLIN}} + if n, err := unix.Poll(pfd, 10000); n != 1 || err != nil || pfd[0].Revents&unix.POLLHUP == 0 { + t.Fatalf("the passed descriptors are still open: poll = %d, %v, revents %#x", n, err, pfd[0].Revents) + } +} + +// A gone output reader must not hold the exit back while the stream that +// would close the remote output is lost behind the settled operation. +func TestGoneReaderDoesNotHoldExit(t *testing.T) { + f := newFixture(t, first(func(c net.Conn) io.ReadWriteCloser { return &holdEvents{Conn: c} })) + r, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + r.Close() + cmd := f.command("sh", "-c", "echo hi") + cmd.Stdout = w + err = cmd.Start() + w.Close() + if err != nil { + t.Fatal(err) + } + wait := make(chan error, 1) + go func() { wait <- cmd.Wait() }() + select { + case err := <-wait: + if err != nil { + t.Fatalf("Wait = %v", err) + } + case <-time.After(10 * time.Second): + cmd.Process.Kill() + t.Fatal("the shim did not exit") + } +} + +// A receiver with SO_PASSCRED sees the relay's own credentials on output: +// its pid and the Session's uid and gid. +func TestUnixSocketOutputNamesTheRelay(t *testing.T) { + f := newFixture(t, nil) + sv, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + if err != nil { + t.Fatal(err) + } + defer unix.Close(sv[0]) + if err := unix.SetsockoptInt(sv[0], unix.SOL_SOCKET, unix.SO_PASSCRED, 1); err != nil { + t.Fatal(err) + } + out := os.NewFile(uintptr(sv[1]), "output") + cmd := f.command("sh", "-c", "echo hi") + cmd.Stdout = out + err = cmd.Start() + out.Close() + if err != nil { + t.Fatal(err) + } + buf, oob := make([]byte, 64), make([]byte, unix.CmsgSpace(unix.SizeofUcred)) + n, oobn, _, _, err := unix.Recvmsg(sv[0], buf, oob, 0) + if err != nil { + t.Fatal(err) + } + if err := cmd.Wait(); err != nil { + t.Fatalf("Wait = %v", err) + } + msgs, err := unix.ParseSocketControlMessage(oob[:oobn]) + if err != nil || len(msgs) != 1 { + t.Fatalf("control messages %v, %v", msgs, err) + } + cred, err := unix.ParseUnixCredentials(&msgs[0]) + if err != nil { + t.Fatal(err) + } + if string(buf[:n]) != "hi\n" || int(cred.Pid) != f.relay.Process.Pid || int(cred.Uid) != f.uid || int(cred.Gid) != f.gid { + t.Fatalf("read %q from %+v; the relay is pid %d, uid %d, gid %d", buf[:n], *cred, f.relay.Process.Pid, f.uid, f.gid) + } +} + +// A write the kernel refuses the Session fails as a write failure, never +// with the daemon's authority: lowering oom_score_adj needs +// CAP_SYS_RESOURCE, which the relay lacks. +func TestOutputWritesHaveTheSessionsAuthority(t *testing.T) { + f := newFixture(t, nil) + before, err := os.ReadFile("/proc/self/oom_score_adj") + if err != nil { + t.Fatal(err) + } + adj, err := os.OpenFile("/proc/self/oom_score_adj", os.O_WRONLY, 0) + if err != nil { + t.Fatal(err) + } + defer adj.Close() + cmd := f.command("bash", "-c", `trap "" PIPE; echo -1000; while echo x; do sleep 0.05; done 2>/dev/null; exit 3`) + cmd.Stdout = adj + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + wait := make(chan error, 1) + go func() { wait <- cmd.Wait() }() + select { + case err := <-wait: + if exitCode(err) != 3 { + t.Fatalf("Wait = %v, want exit 3 after the remote output closed", err) + } + case <-time.After(10 * time.Second): + cmd.Process.Kill() + t.Fatal("the failed write never closed the remote output") + } + if after, err := os.ReadFile("/proc/self/oom_score_adj"); err != nil || !bytes.Equal(after, before) { + t.Fatalf("oom_score_adj %q, %v; was %q", after, err, before) + } +} + +// Output that nobody reads holds back only its own invocation: another +// invocation starts, reads stdin, gets a signal and exits, and the blocked +// one still gets its signal and its cancel. +func TestBlockedOutputHoldsOnlyItsInvocation(t *testing.T) { + f := newFixture(t, nil) + r, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + defer r.Close() + a := f.command("bash", "-c", `echo $$ > a.pid; trap 'touch a.int' INT; while :; do yes; done`) + a.Stdout = w + err = a.Start() + w.Close() + if err != nil { + t.Fatal(err) + } + defer a.Process.Kill() + size, err := unix.FcntlInt(r.Fd(), unix.F_GETPIPE_SZ, 0) + if err != nil { + t.Fatal(err) + } + await(t, "the first invocation's pipe to fill", func() bool { + n, err := unix.IoctlGetInt(int(r.Fd()), unix.TIOCINQ) // FIONREAD + return err == nil && n == size + }) + + b := f.command("bash", "-c", `trap 'exit 6' TERM; read line; echo "got $line"; while :; do sleep 0.05; done`) + b.Stdin = strings.NewReader("hello\n") + startLine(t, b, "got hello") + if err := b.Process.Signal(syscall.SIGTERM); err != nil { + t.Fatal(err) + } + if err := b.Wait(); exitCode(err) != 6 { + t.Fatalf("second Wait = %v, want exit 6", err) + } + + if err := a.Process.Signal(syscall.SIGINT); err != nil { + t.Fatal(err) + } + await(t, "the blocked invocation's trap", func() bool { + _, err := os.Stat(filepath.Join(f.dir, "a.int")) + return err == nil + }) + pid, err := os.ReadFile(filepath.Join(f.dir, "a.pid")) + if err != nil { + t.Fatal(err) + } + remote := "/proc/" + strings.TrimSpace(string(pid)) + a.Process.Kill() + a.Wait() + await(t, "the blocked invocation's cancel", func() bool { + _, err := os.Stat(remote) + return errors.Is(err, os.ErrNotExist) + }) +} + +// The broker never takes a descriptor from the relay: one sent with a +// message is closed as the broker reads it. +func TestRelayDescriptorsAreDiscarded(t *testing.T) { + sv, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + if err != nil { + t.Fatal(err) + } + end := os.NewFile(uintptr(sv[0]), "broker") + b, err := Start(Config{ + Relay: end, + Scope: sp.ScopePOSIXSession, + Dial: func(context.Context) (io.ReadWriteCloser, error) { return nil, errors.New("no sandbox") }, + }) + end.Close() + if err != nil { + unix.Close(sv[1]) + t.Fatal(err) + } + defer b.Close() + rf := os.NewFile(uintptr(sv[1]), "relay") + c, err := net.FileConn(rf) + rf.Close() + if err != nil { + t.Fatal(err) + } + defer func() { + b.Close() + c.Close() + }() + var p [2]int + if err := unix.Pipe2(p[:], unix.O_CLOEXEC); err != nil { + t.Fatal(err) + } + defer unix.Close(p[0]) + var open bytes.Buffer + req := processshim.Request{Version: processshim.Version, ExecPath: []byte("/bin/undeclared"), Argv: [][]byte{[]byte("undeclared")}, Cwd: []byte("/")} + if err := sandboxwire.WriteFrame(&open, processshim.Frame(processshim.Open{ID: 1, Request: req})); err != nil { + t.Fatal(err) + } + _, _, err = c.(*net.UnixConn).WriteMsgUnix(open.Bytes(), unix.UnixRights(p[1]), nil) + unix.Close(p[1]) + if err != nil { + t.Fatal(err) + } + c.SetReadDeadline(time.Now().Add(10 * time.Second)) + for { + fr, err := sandboxwire.ReadFrame(c, processshim.MaxFrameBytes) + if err != nil { + t.Fatal(err) + } + m, err := processshim.DecodeBroker(fr) + if _, ok := m.(processshim.StopInput); ok { + continue + } + if exit, ok := m.(processshim.Exit); err != nil || !ok || exit.Result.Code != processshim.ExitNotFound { + t.Fatalf("broker answered %#v, %v", m, err) + } + break + } + pfd := []unix.PollFd{{Fd: int32(p[0]), Events: unix.POLLIN}} + if n, err := unix.Poll(pfd, 5000); n != 1 || err != nil || pfd[0].Revents&unix.POLLHUP == 0 { + t.Fatalf("the broker holds the relay's descriptor: poll = %d, %v, revents %#x", n, err, pfd[0].Revents) + } + if err := b.Err(); err != nil { + t.Fatalf("Err = %v", err) + } +} + +// A Busy acknowledgement is retried, or output past the replay limit would +// never arrive. +func TestBusyAckIsRetried(t *testing.T) { + f := newFixture(t, first(func(c net.Conn) io.ReadWriteCloser { return &busyAck{Conn: c} })) + cmd := f.command("seq", "1", "2000000") + var out countWriter + cmd.Stdout = &out + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + wait := make(chan error, 1) + go func() { wait <- cmd.Wait() }() + select { + case err := <-wait: + if err != nil || out.n != 14888896 { + t.Fatalf("Wait = %v after %d bytes", err, out.n) + } + case <-time.After(30 * time.Second): + cmd.Process.Kill() + t.Fatal("output stalled after a Busy acknowledgement") + } +} + +// Concurrent invocations on one terminal leave it as it was: the first saves +// its mode and the last restores it. +func TestSharedTerminalRestoredByLastUser(t *testing.T) { + f := newFixture(t, nil) + ptm, pts, err := pty.Open() + if err != nil { + t.Fatal(err) + } + defer ptm.Close() + defer pts.Close() + go io.Copy(io.Discard, ptm) + mode := func() *unix.Termios { + tio, err := unix.IoctlGetTermios(int(pts.Fd()), unix.TCGETS) + if err != nil { + t.Fatal(err) + } + return tio + } + raw := func() bool { return mode().Lflag&unix.ICANON == 0 } + file := func(name string) string { return filepath.Join(f.dir, name) } + exists := func(name string) func() bool { + return func() bool { _, err := os.Stat(file(name)); return err == nil } + } + before := mode() + run := func(script string) *exec.Cmd { + cmd := f.command("bash", "-c", script) + cmd.Stdin, cmd.Stdout, cmd.Stderr = pts, pts, pts + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + return cmd + } + a := run("until [ -e a.done ]; do sleep 0.05; done") + await(t, "the first invocation to make the terminal raw", raw) + b := run("touch b.up; until [ -e b.done ]; do sleep 0.05; done") + await(t, "the second invocation to run", exists("b.up")) + os.WriteFile(file("a.done"), nil, 0o644) + if err := a.Wait(); err != nil { + t.Fatalf("first Wait = %v", err) + } + if !raw() { + t.Fatal("the first invocation to leave restored the terminal under the second") + } + os.WriteFile(file("b.done"), nil, 0o644) + if err := b.Wait(); err != nil { + t.Fatalf("second Wait = %v", err) + } + if after := mode(); *after != *before { + t.Fatalf("terminal mode not restored:\nbefore %+v\nafter %+v", *before, *after) + } +} + +// A Busy stdin write is repeated from the first byte the service did not +// accept, so the program reads every byte once. +func TestBusyStdinIsRetried(t *testing.T) { + var writes atomic.Int32 + f := newFixture(t, first(func(c net.Conn) io.ReadWriteCloser { + return intercept(c, func(fr sandboxwire.Frame) verdict { + if fr.Type == sp.OpWriteStdin && writes.Add(1) == 2 { + return refuseBusy + } + return pass + }) + })) + var in []byte + for i := range 30000 { + in = strconv.AppendInt(in, int64(i), 10) + in = append(in, '\n') + } + cmd := f.command("sh", "-c", "cat") + cmd.Stdin = bytes.NewReader(in) + wait := make(chan error, 1) + var out []byte + go func() { + var err error + out, err = cmd.Output() + wait <- err + }() + select { + case err := <-wait: + if err != nil || !bytes.Equal(out, in) || writes.Load() < 2 { + t.Fatalf("Output = %v after %d stdin writes; read %d bytes of %d, equal %v", err, writes.Load(), len(out), len(in), bytes.Equal(out, in)) + } + case <-time.After(30 * time.Second): + cmd.Process.Kill() + t.Fatal("stdin stalled after a Busy write") + } +} + +// A Cancel whose response is lost with its stream is not sent again: the +// scope gets one TERM, and KILL after the grace. +func TestUncertainCancelIsNotReplayed(t *testing.T) { + var cancels atomic.Int32 + f := newFixture(t, func(_ int32, c net.Conn) io.ReadWriteCloser { + return intercept(c, func(fr sandboxwire.Frame) verdict { + if fr.Type == sp.OpCancel && cancels.Add(1) == 1 { + return loseResponse + } + return pass + }) + }) + r, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + defer r.Close() + cmd := f.command("sh", "-c", "trap '' TERM; echo up; sleep 1000") + cmd.Stdout = w + err = cmd.Start() + w.Close() + if err != nil { + t.Fatal(err) + } + out := bufio.NewReader(r) + if line, err := out.ReadString('\n'); line != "up\n" { + t.Fatalf("first line %q, %v", line, err) + } + cmd.Process.Kill() + cmd.Wait() + read := make(chan error, 1) + go func() { + _, err := io.ReadAll(out) + read <- err + }() + select { + case <-read: + case <-time.After(10 * time.Second): + t.Fatal("the cancelled program still runs") + } + if n := cancels.Load(); n != 1 { + t.Fatalf("%d Cancel requests, want 1", n) + } +} + +// Losing the shim while the first stream is still connecting ends the +// invocation: the broker closes the shim's descriptors. +func TestShimLossEndsWaitForStream(t *testing.T) { + dialed := make(chan struct{}) + f := newFixture(t, func(n int32, c net.Conn) io.ReadWriteCloser { + if n > 1 { + return c + } + c.Close() + close(dialed) + // A service that never answers. + broker, svc := net.Pipe() + go io.Copy(io.Discard, svc) + return broker + }) + r, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + defer r.Close() + cmd := f.command("sh", "-c", "echo started") + cmd.Stdout = w + err = cmd.Start() + w.Close() + if err != nil { + t.Fatal(err) + } + select { + case <-dialed: + case <-time.After(10 * time.Second): + cmd.Process.Kill() + t.Fatal("the broker never dialed") + } + cmd.Process.Kill() + cmd.Wait() + read := make(chan []byte, 1) + go func() { + data, _ := io.ReadAll(r) + read <- data + }() + select { + case data := <-read: + if len(data) != 0 { + t.Fatalf("read %q", data) + } + case <-time.After(10 * time.Second): + t.Fatal("the broker still holds the lost shim's descriptors") + } +} + +// A Start retried after its response was lost finds the operation still +// starting. The program runs, and gets stdin, only once Started arrives: it +// reads all of its input and then end of file. A StartFailed event instead +// is the typed start failure. +func TestRetriedStartWaitsForStarted(t *testing.T) { + for _, fails := range []bool{false, true} { + t.Run(fmt.Sprintf("start fails %v", fails), func(t *testing.T) { + svc := newFakeService() + svc.startLate, svc.startFails = true, fails + f := newFixture(t, svc.serve(t.Context(), func(fr sandboxwire.Frame) bool { return fr.Type == sp.OpStart })) + // Once the retried Start's Attach found the operation starting, + // it starts when the broker observes it without forwarding stdin. + // A broker that forwards stdin first gets the refusal it counts. + go func() { + select { + case <-svc.attachedStarting: + case <-time.After(10 * time.Second): + return // runFor fails the test + } + for deadline := time.Now().Add(10 * time.Second); !awaitsStarted(); time.Sleep(time.Millisecond) { + if time.Now().After(deadline) { + return + } + } + svc.runStarting() + }() + const in = "stdin for a program that was still starting\n" + out, stderr, err := runFor(t, f.command("sh", "-c", "cat"), in) + refused := svc.counts().refusedIn + switch { + case fails && (exitCode(err) != processshim.ExitNotFound || !strings.Contains(stderr, "no such file")): + t.Fatalf("Run = %v; stderr %q", err, stderr) + case !fails && (err != nil || out != in || refused != 0): + t.Fatalf("Run = %v after %d stdin requests refused while starting; stdout %q, stderr %q", err, refused, out, stderr) + } + }) + } +} + +// awaitsStarted reports whether an invocation observes its operation with no +// stdin pump, as one does until Started arrives. +func awaitsStarted() bool { + buf := make([]byte, 1<<16) + for { + n := runtime.Stack(buf, true) + if n < len(buf) { + buf = buf[:n] + break + } + buf = make([]byte, 2*len(buf)) + } + stacks := string(buf) + return strings.Contains(stacks, "processbroker.(*invocation).observe(") && + !strings.Contains(stacks, "processbroker.(*invocation).pumpStdin(") +} + +// A stdin write lost with its stream after the service took part of it +// resumes from the offset Inspect reports: the program reads every byte +// once, and stdin closes at the end of the input. +func TestUncertainStdinWriteResumes(t *testing.T) { + svc := newFakeService() + svc.partialFirst = true + f := newFixture(t, svc.serve(t.Context(), func(fr sandboxwire.Frame) bool { return fr.Type == sp.OpWriteStdin })) + const in = "stdin that crosses a lost stream\n" + out, stderr, err := runFor(t, f.command("sh", "-c", "cat"), in) + c := svc.counts() + if err != nil || out != in || string(c.stdin) != in || c.closedAt != uint64(len(in)) || c.writes < 2 { + t.Fatalf("Run = %v after %d writes; stdout %q, stderr %q; the service took %q and closed stdin at %d", err, c.writes, out, stderr, c.stdin, c.closedAt) + } +} + +// When the accepted stdin offset cannot be learned after a lost write, the +// shim exits with 255 and the reason, and the program is cancelled, rather +// than both waiting for input that never comes. +func TestUnresolvedStdinEndsTheInvocation(t *testing.T) { + svc := newFakeService() + svc.partialFirst, svc.inspectFails = true, true + f := newFixture(t, svc.serve(t.Context(), func(fr sandboxwire.Frame) bool { return fr.Type == sp.OpWriteStdin })) + _, stderr, err := runFor(t, f.command("sh", "-c", "cat"), "stdin\n") + if exitCode(err) != processshim.ExitLost || !strings.Contains(stderr, "stdin could not be resumed: inspect failed") { + t.Fatalf("Run = %v; stderr %q", err, stderr) + } + await(t, "the program's cancel", func() bool { return svc.counts().cancels == 1 }) +} + +// A background process outlives the leader, which exits as it takes a stdin +// write whose response is lost. Inspect shows the leader exited, perhaps +// before its Exited event arrives: stdin closes after the bytes the service +// took, the background process gets end of file, and the operation settles. +func TestBackgroundReaderGetsEOFAfterUncertainWrite(t *testing.T) { + svc := newFakeService() + svc.leaderExits = true + f := newFixture(t, svc.serve(t.Context(), func(fr sandboxwire.Frame) bool { return fr.Type == sp.OpWriteStdin })) + const in = "stdin for a background reader\n" + out, stderr, err := runFor(t, f.command("sh", "-c", "cat"), in) + if err != nil || out != in { + t.Fatalf("Run = %v; stdout %q, stderr %q", err, out, stderr) + } + await(t, "the operation's release", func() bool { return svc.counts().released }) + if c := svc.counts(); c.closedAt != uint64(len(in)) { + t.Fatalf("stdin closed at %d, not %d", c.closedAt, len(in)) + } +} + +// runFor runs cmd with stdin in and returns its stdout, stderr and error. It +// fails the test when cmd still runs after 10s; output the relay still +// holds open 5s after cmd exits is an exec.ErrWaitDelay error. +func runFor(t *testing.T, cmd *exec.Cmd, in string) (string, string, error) { + t.Helper() + var stdout, stderr bytes.Buffer + cmd.Stdin, cmd.Stdout, cmd.Stderr = strings.NewReader(in), &stdout, &stderr + cmd.WaitDelay = 5 * time.Second + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + done := make(chan error, 1) + go func() { done <- cmd.Wait() }() + select { + case err := <-done: + return stdout.String(), stderr.String(), err + case <-time.After(10 * time.Second): + cmd.Process.Kill() + <-done + t.Fatalf("%s still ran after 10s; stdout %q, stderr %q", cmd.Args, stdout.String(), stderr.String()) + return "", "", nil + } +} + +type fixture struct { + dir, bin string + uid, gid int // the relay's + service string + wrap func(int32, net.Conn) io.ReadWriteCloser + dials atomic.Int32 + relay *exec.Cmd + broker *Broker +} + +// newFixture starts a relay and its broker, whose shims are this binary +// under the names bash, sh, env and seq. As root the relay runs as viewID. A +// non-nil wrap wraps each stream to the service, which it gets with the +// stream's number, counting from 1. +func newFixture(t *testing.T, wrap func(int32, net.Conn) io.ReadWriteCloser) *fixture { + t.Helper() + self, err := os.Executable() + if err != nil { + t.Fatal(err) + } + dir, err := os.MkdirTemp("", "processbroker-") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { os.RemoveAll(dir) }) + if err := os.Chmod(dir, 0o755); err != nil { + t.Fatal(err) + } + f := &fixture{dir: dir, uid: os.Getuid(), gid: os.Getgid(), service: service.socket(t), wrap: wrap} + if f.uid == 0 { + f.uid, f.gid = viewID, viewID + } + f.bin = filepath.Join(f.dir, "bin") + if err := os.Mkdir(f.bin, 0o755); err != nil { + t.Fatal(err) + } + paths := map[string]string{} + for _, name := range []string{"bash", "sh", "env", "seq"} { + local := filepath.Join(f.bin, name) + if err := os.Symlink(self, local); err != nil { + t.Fatal(err) + } + paths[local] = name + } + end := f.startRelay(t, self) + b, err := Start(Config{ + Relay: end, + Executables: Executables{Paths: paths}, + Environment: Environment{ + Pass: []string{"KEEP", "LEAK"}, + Sandbox: map[string]string{"PATH": "/usr/bin:/bin", "HOME": "/home/sandbox", "LANG": "C.UTF-8"}, + Tool: map[string]string{"TOOL": "1"}, + }, + Scope: sp.ScopePOSIXSession, + Dial: f.dial, + CancelGrace: time.Second, + }) + end.Close() + if err != nil { + f.relay.Process.Kill() + f.relay.Wait() + t.Fatal(err) + } + f.broker = b + t.Cleanup(func() { + b.Close() + f.waitRelay(t) + }) + return f +} + +// startRelay starts this binary as the relay, listening at the shim socket +// in f.dir, and returns the broker's end of its connection. +func (f *fixture) startRelay(t *testing.T, self string) *os.File { + t.Helper() + sv, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + if err != nil { + t.Fatal(err) + } + end, relayEnd := os.NewFile(uintptr(sv[0]), "broker"), os.NewFile(uintptr(sv[1]), "relay") + defer relayEnd.Close() + sock := filepath.Join(f.dir, processshim.SocketName) + ln, err := net.ListenUnix("unix", &net.UnixAddr{Name: sock, Net: "unix"}) + if err != nil { + end.Close() + t.Fatal(err) + } + ln.SetUnlinkOnClose(false) + lf, err := ln.File() + ln.Close() + if err == nil { + defer lf.Close() + err = os.Chmod(sock, 0o666) + } + if err != nil { + end.Close() + t.Fatal(err) + } + f.relay = exec.Command(self) + f.relay.Env = []string{relayEnv + "=1", "GORACE=atexit_sleep_ms=0"} + f.relay.ExtraFiles = []*os.File{relayEnd, lf} // processshim.RelayBrokerFD and RelayListenerFD + f.relay.Stderr = os.Stderr + if f.uid != os.Getuid() { + f.relay.SysProcAttr = &syscall.SysProcAttr{Credential: &syscall.Credential{Uid: uint32(f.uid), Gid: uint32(f.gid)}} + } + if err := f.relay.Start(); err != nil { + end.Close() + t.Fatal(err) + } + return end +} + +// waitRelay waits for the relay, which exits once the broker's connection +// ends. +func (f *fixture) waitRelay(t *testing.T) { + done := make(chan error, 1) + go func() { done <- f.relay.Wait() }() + select { + case err := <-done: + if err != nil { + t.Errorf("relay: %v", err) + } + case <-time.After(10 * time.Second): + f.relay.Process.Kill() + <-done + t.Error("the relay still runs 10s after the broker closed") + } +} + +func (f *fixture) command(name string, args ...string) *exec.Cmd { + cmd := exec.Command(filepath.Join(f.bin, name), args...) + // A race-enabled shim would otherwise sleep a second before exiting. + cmd.Env = []string{shimSocketEnv + "=" + filepath.Join(f.dir, processshim.SocketName), "GORACE=atexit_sleep_ms=0"} + cmd.Dir = f.dir + return cmd +} + +func (f *fixture) dial(ctx context.Context) (io.ReadWriteCloser, error) { + c, err := new(net.Dialer).DialContext(ctx, "unix", f.service) + if err != nil { + return nil, err + } + if n := f.dials.Add(1); f.wrap != nil { + return f.wrap(n, c), nil + } + return c, nil +} + +// first applies wrap to the first stream alone. +func first(wrap func(net.Conn) io.ReadWriteCloser) func(int32, net.Conn) io.ReadWriteCloser { + return func(n int32, c net.Conn) io.ReadWriteCloser { + if n == 1 { + return wrap(c) + } + return c + } +} + +// verdict is what intercept does with a request from the broker. +type verdict int + +const ( + pass verdict = iota // pass it to the service + refuseBusy // answer Busy with no effect, as a service at its request limit does + loseResponse // pass it on, then lose the stream in place of its response +) + +// intercept relays a stream between the broker and the service and applies +// decide to each request. +func intercept(svc net.Conn, decide func(sandboxwire.Frame) verdict) io.ReadWriteCloser { + broker, relay := net.Pipe() + var mu sync.Mutex // writes to the broker + var lost atomic.Uint64 + cut := func() { + svc.Close() + relay.Close() + } + go func() { + defer cut() + for { + fr, err := sandboxwire.ReadFrame(svc, sandboxwire.MaxPayload) + if err != nil || fr.RequestID != 0 && fr.RequestID == lost.Load() { + return + } + mu.Lock() + err = sandboxwire.WriteFrame(relay, fr) + mu.Unlock() + if err != nil { + return + } + } + }() + go func() { + defer cut() + for { + fr, err := sandboxwire.ReadFrame(relay, sandboxwire.MaxPayload) + if err != nil { + return + } + switch decide(fr) { + case refuseBusy: + m := sp.ResponseFailure{Request: fr.Type, Failure: *sp.Fail(sp.CodeBusy, sandboxwire.EffectNone, "busy")} + mu.Lock() + err = sandboxwire.WriteFrame(relay, sandboxwire.Frame{Type: m.MessageType(), RequestID: fr.RequestID, Payload: sp.Encode(m)}) + mu.Unlock() + case loseResponse: + lost.Store(fr.RequestID) + fallthrough + default: + err = sandboxwire.WriteFrame(svc, fr) + } + if err != nil { + return + } + } + }() + return broker +} + +// cutConn loses the link after left bytes, mid-frame if it falls there. +type cutConn struct { + net.Conn + left int +} + +func (c *cutConn) Read(p []byte) (int, error) { + if c.left <= 0 { + c.Conn.Close() + return 0, net.ErrClosed + } + n, err := c.Conn.Read(p[:min(len(p), c.left)]) + c.left -= n + return n, err +} + +// holdEvents holds the service's events until the scope closes, then passes +// them on at once and loses the link before anything that follows. +type holdEvents struct { + net.Conn + held, out bytes.Buffer + cut bool +} + +func (c *holdEvents) Read(p []byte) (int, error) { + for c.out.Len() == 0 { + if c.cut { + c.Conn.Close() + return 0, net.ErrClosed + } + fr, err := sandboxwire.ReadFrame(c.Conn, sandboxwire.MaxPayload) + if err != nil { + return 0, err + } + dst := &c.out + if fr.RequestID == 0 { + dst = &c.held + } + if err := sandboxwire.WriteFrame(dst, fr); err != nil { + return 0, err + } + if fr.Type == sp.EventScopeClosed { + c.held.WriteTo(&c.out) + c.cut = true + } + } + return c.out.Read(p) +} + +// busyAck answers the first acknowledgement with Busy after the service +// applied it. +type busyAck struct { + net.Conn + out bytes.Buffer + sent bool +} + +func (c *busyAck) Read(p []byte) (int, error) { + for c.out.Len() == 0 { + fr, err := sandboxwire.ReadFrame(c.Conn, sandboxwire.MaxPayload) + if err != nil { + return 0, err + } + if !c.sent && fr.Type == sandboxwire.ResponseType(sp.OpAckEvents) { + c.sent = true + fr.Payload = sp.Encode(sp.ResponseFailure{Request: sp.OpAckEvents, Failure: *sp.Fail(sp.CodeBusy, sandboxwire.EffectNone, "busy")}) + } + if err := sandboxwire.WriteFrame(&c.out, fr); err != nil { + return 0, err + } + } + return c.out.Read(p) +} + +type countWriter struct{ n int } + +func (w *countWriter) Write(p []byte) (int, error) { + w.n += len(p) + return len(p), nil +} + +func start(t *testing.T, cmd *exec.Cmd) io.Reader { + t.Helper() + out, err := cmd.StdoutPipe() + if err != nil { + t.Fatal(err) + } + if err := cmd.Start(); err != nil { + t.Fatal(err) + } + return out +} + +// startLine starts cmd and reads its first line, which must be want. +func startLine(t *testing.T, cmd *exec.Cmd, want string) io.Reader { + t.Helper() + out := bufio.NewReader(start(t, cmd)) + if line, err := out.ReadString('\n'); err != nil || line != want+"\n" { + t.Fatalf("first line %q, %v", line, err) + } + return out +} + +// await polls done until it holds, for at most 10s. +func await(t *testing.T, what string, done func() bool) { + t.Helper() + for deadline := time.Now().Add(10 * time.Second); !done(); time.Sleep(10 * time.Millisecond) { + if time.Now().After(deadline) { + t.Fatalf("timed out waiting for %s", what) + } + } +} + +func exitCode(err error) int { + var ee *exec.ExitError + if errors.As(err, &ee) { + return ee.ExitCode() + } + if err == nil { + return 0 + } + return -1 +} + +// processService is the process service the tests share. +type processService struct { + once sync.Once + dir string + sock string + cmd *exec.Cmd + err error +} + +var service processService + +func (s *processService) socket(t *testing.T) string { + t.Helper() + s.once.Do(func() { s.err = s.start() }) + if s.err != nil { + t.Fatalf("process service: %v", s.err) + } + return s.sock +} + +func (s *processService) start() error { + dir, err := os.MkdirTemp("", "processbroker") + if err != nil { + return err + } + s.dir = dir + bin := os.Getenv(serviceEnv) + if bin == "" { + bin = filepath.Join(dir, "processserve") + if err := build(bin, "../../../sandboxio/testdata/processserve"); err != nil { + return err + } + } + s.sock = filepath.Join(dir, "service.sock") + s.cmd = exec.Command(bin, s.sock) + s.cmd.Stderr = os.Stderr + out, err := s.cmd.StdoutPipe() + if err != nil { + return err + } + if err := s.cmd.Start(); err != nil { + return err + } + if line, err := bufio.NewReader(out).ReadString('\n'); line != "ready\n" { + return fmt.Errorf("service said %q: %v", line, err) + } + return nil +} + +func (s *processService) stop() { + if s.cmd != nil && s.cmd.Process != nil { + s.cmd.Process.Kill() + s.cmd.Wait() + } + if s.dir != "" { + os.RemoveAll(s.dir) + } +} + +func build(out, pkg string) error { + cmd := exec.Command("go", "build", "-o", out, pkg) + cmd.Env = append(os.Environ(), "CGO_ENABLED=0") + if msg, err := cmd.CombinedOutput(); err != nil { + return fmt.Errorf("go build %s: %v\n%s", pkg, err, msg) + } + return nil +} diff --git a/apps/daemon/internal/processbroker/broker_other.go b/apps/daemon/internal/processbroker/broker_other.go new file mode 100644 index 00000000..7c4971c8 --- /dev/null +++ b/apps/daemon/internal/processbroker/broker_other.go @@ -0,0 +1,23 @@ +//go:build !linux + +package processbroker + +import "errors" + +// ErrUnsupported is returned on platforms without the process broker. +var ErrUnsupported = errors.New("processbroker: unsupported on this platform") + +// Broker is unavailable on this platform. +type Broker struct{} + +// Start returns ErrUnsupported. +func Start(Config) (*Broker, error) { return nil, ErrUnsupported } + +// Close does nothing. +func (*Broker) Close() error { return nil } + +// Done returns nil. +func (*Broker) Done() <-chan struct{} { return nil } + +// Err returns ErrUnsupported. +func (*Broker) Err() error { return ErrUnsupported } diff --git a/apps/daemon/internal/processbroker/config.go b/apps/daemon/internal/processbroker/config.go new file mode 100644 index 00000000..580cd032 --- /dev/null +++ b/apps/daemon/internal/processbroker/config.go @@ -0,0 +1,183 @@ +package processbroker + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "log/slog" + "maps" + "os" + "path" + "slices" + "strings" + "time" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/agent" + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" +) + +// Config configures one Session's broker. +type Config struct { + // Relay is the broker's end of the connection to the Session's process + // relay, sessionview's View.Relay. The broker uses a duplicate; the + // caller keeps Relay. + Relay *os.File + // Executables is the declared executable table. + Executables Executables + // Environment is the declared environment policy. + Environment Environment + // Scope is the containment each operation starts in. The service must + // declare it. + Scope sp.Scope + // Dial opens a Process stream to the Session's sandbox. + Dial func(context.Context) (io.ReadWriteCloser, error) + // CancelGrace is the Cancel grace when a shim is lost before its program + // exits. The service caps it at its CancelGraceLimitMillis. + CancelGrace time.Duration + // Logger receives the broker's decisions; nil uses slog.Default. + Logger *slog.Logger +} + +// Executables maps local invocation paths to remote executables. A remote +// executable is a bare name, resolved on the remote PATH, or an absolute +// sandbox path. +type Executables struct { + // Names maps a name in the view's shim directory to its remote executable. + Names map[string]string + // Paths maps an absolute view path the shim is bound over, such as + // /bin/bash, to its remote executable. + Paths map[string]string +} + +// Environment is the remote environment policy. A name set in more than one +// place takes the value from the later of Pass, Sandbox and Tool. Nothing +// else reaches the sandbox. +type Environment struct { + // Pass names the shim environment entries that pass through. + Pass []string + // Sandbox holds fixed sandbox values such as HOME, PATH, TMPDIR and LANG. + Sandbox map[string]string + // Tool is the Environment's tool environment. + Tool map[string]string +} + +var ( + // ErrInvalidConfig wraps every configuration error. + ErrInvalidConfig = errors.New("processbroker: invalid configuration") + // ErrRelayLost wraps the end of the relay connection before Close, or a + // relay message that breaks the IPC. The broker then serves nothing + // more, and the Session fails. + ErrRelayLost = errors.New("processbroker: process relay lost") +) + +func (c *Config) validate() error { + invalid := func(format string, args ...any) error { + return fmt.Errorf("%w: %s", ErrInvalidConfig, fmt.Sprintf(format, args...)) + } + switch { + case c.Relay == nil: + return invalid("no relay connection") + case !c.Scope.Valid(): + return invalid("scope %d", c.Scope) + case c.Dial == nil: + return invalid("no dial function") + case c.CancelGrace < 0: + return invalid("negative cancel grace") + } + names := slices.Concat(c.Environment.Pass, slices.Collect(maps.Keys(c.Environment.Sandbox)), slices.Collect(maps.Keys(c.Environment.Tool))) + for _, name := range names { + if !validEnvName(name) { + return invalid("environment name %q", name) + } + } + _, path1 := c.Environment.Sandbox["PATH"] + _, path2 := c.Environment.Tool["PATH"] + for name, remote := range c.Executables.Names { + if name == "" || name == "." || name == ".." || strings.ContainsAny(name, "/\x00") { + return invalid("executable name %q", name) + } + if err := checkRemote(remote, path1 || path2); err != nil { + return invalid("executable %q: %v", name, err) + } + } + for local, remote := range c.Executables.Paths { + if !path.IsAbs(local) || path.Clean(local) != local || underPrivate(local) || strings.Contains(local, "\x00") { + return invalid("executable path %q", local) + } + if err := checkRemote(remote, path1 || path2); err != nil { + return invalid("executable %q: %v", local, err) + } + } + return nil +} + +func checkRemote(remote string, havePATH bool) error { + switch { + case remote == "" || strings.Contains(remote, "\x00"): + return errors.New("empty remote executable") + case strings.Contains(remote, "/") && !path.IsAbs(remote): + return fmt.Errorf("remote executable %q is neither a name nor absolute", remote) + case !strings.Contains(remote, "/") && !havePATH: + return fmt.Errorf("remote name %q needs PATH in the sandbox or tool environment", remote) + } + return nil +} + +func validEnvName(name string) bool { + return name != "" && !strings.ContainsAny(name, "=\x00") +} + +func underPrivate(p string) bool { + return p == agent.ViewPrivateRoot || strings.HasPrefix(p, agent.ViewPrivateRoot+"/") +} + +// resolve returns the remote executable for the shim's exec path. A relative +// path resolves against cwd; resolution is lexical, because the view's +// symlinks are not the broker's to follow. +func (x Executables) resolve(execPath, cwd string) (string, bool) { + p := execPath + if !path.IsAbs(p) { + p = path.Join(cwd, p) + } + p = path.Clean(p) + if name, ok := strings.CutPrefix(p, agent.ViewPrivateRoot+"/"+agent.ViewShimName+"/"); ok { + remote, ok := x.Names[name] + return remote, ok + } + remote, ok := x.Paths[p] + return remote, ok +} + +// privateMarker is the view prefix no value may carry into the sandbox. +var privateMarker = []byte(agent.ViewPrivateRoot) + +// compose builds the remote environment from the shim's environ. It returns +// the names it dropped because their value names the private directory. +func (e Environment) compose(environ [][]byte) (env []sp.EnvVar, dropped []string) { + values := map[string][]byte{} + for _, entry := range environ { + name, value, ok := bytes.Cut(entry, []byte("=")) + if !ok || !slices.Contains(e.Pass, string(name)) { + continue + } + if _, seen := values[string(name)]; !seen { // getenv returns the first + values[string(name)] = value + } + } + for name, value := range e.Sandbox { + values[name] = []byte(value) + } + for name, value := range e.Tool { + values[name] = []byte(value) + } + for _, name := range slices.Sorted(maps.Keys(values)) { + if bytes.Contains(values[name], privateMarker) { + dropped = append(dropped, name) + continue + } + env = append(env, sp.EnvVar{Name: []byte(name), Value: values[name]}) + } + return env, dropped +} diff --git a/apps/daemon/internal/processbroker/doc.go b/apps/daemon/internal/processbroker/doc.go new file mode 100644 index 00000000..a511ea96 --- /dev/null +++ b/apps/daemon/internal/processbroker/doc.go @@ -0,0 +1,107 @@ +// Package processbroker runs a Session's shim invocations in its sandbox. +// +// A Harness in the Session view executes oac-process-shim +// (apps/daemon/internal/processshim), which hands its invocation and its +// descriptors 0, 1 and 2 to the Session's process relay. The relay is the +// same binary in relay mode, which sessionview starts in the view as the +// Session user, with no capabilities, before the Harness. It is the only +// process that holds a descriptor from the view: it does every read, write +// and terminal ioctl on them with the Session's own authority. It hands the +// broker the request and the terminal's mode over a socketpair that +// sessionview creates. The broker does all process protocol work: it +// resolves the invocation against the declared executable table and +// environment policy, starts the program with the process protocol +// (internal/sandboxprocess), sends the relay the program's output and asks +// it for stdin, forwards the signals the shim reports, and has the relay +// send the shim the remote exit. Neither the shim nor the relay holds +// credentials. +// +// The broker treats the relay as untrusted Session input. It reads the +// socketpair without a control buffer, so the kernel closes any descriptor +// sent with a message, checks every message against the IPC's types and +// limits, and stops serving at the first message that breaks the IPC. That, +// or the relay's loss, fails the Session: Done closes and Err wraps +// ErrRelayLost. The broker never restarts the relay and never replays output +// whose delivery is uncertain. Close never waits on the relay. +// +// The relay pumps each descriptor on its own, and its dispatch never waits +// on one, so a descriptor nobody reads holds back only its own stream. The +// broker sends the relay at most processshim.OutputWindow bytes of a stream +// that the relay has not reported written, and acknowledges output to the +// process service only after that report, so the service in turn holds back +// the program. +// +// The relay never changes the flags of a passed descriptor, whose open file +// description the Harness shares. It reopens a pipe, FIFO or pty slave +// through /proc/self/fd as its own non-blocking description, uses a socket +// with MSG_DONTWAIT and a regular file or block device as it is, and polls +// any other descriptor before each call, so it waits on a peer only in a +// poll that ending the invocation interrupts. Output on an AF_UNIX socket +// carries the relay's own credentials. +// +// Invocations on one terminal share its saved mode in the relay: the first +// saves it, each runs the terminal raw, and the last to finish restores it. +// +// The broker forwards the exit without waiting for output, and the relay +// answers the shim once the output the program wrote before exiting is +// written. Each output descriptor closes after its stream's last byte is +// written, so a remote background job that keeps its output open keeps the +// Harness's pipe open while the shim still exits when the leader does. +// After the leader exits, remote background processes stay with the Session +// and are not cancelled; a shim lost before the exit cancels the operation's +// scope. +// +// Qualification limits. These behaviors differ from a native child: +// - Stop and continue job control is incomplete. The shim reports TSTP, +// TTIN and TTOU to the remote process group but does not stop itself, so +// the Harness never sees the job stop. SIGSTOP of the shim stops only the +// shim. +// - Stdin is read ahead. The relay reads the shared stdin as the service +// accepts it, so bytes the program never consumes are still taken from a +// stdin the Harness shares with later commands. Forwarding stops when the +// leader exits; remote background readers then see end of file. +// - On a terminal, stderr is merged into the terminal output, as the +// remote PTY merges it. +// - A descriptor 0, 1 or 2 that was closed when the shim started is +// /dev/null, because the Go runtime opens it. +// - The argument list and environment together are limited to +// processshim.MaxRequestBytes, below the kernel's limit. +// - The invocation path is matched lexically: a path reached through a +// symlink the table does not declare fails with 127. +// - After a stdin write whose outcome is uncertain, the broker asks the +// service how much stdin it accepted and sends only the rest. When the +// service cannot say, the shim exits with 255 and the program is +// cancelled. +// - A Cancel whose outcome is uncertain after a lost stream is not sent +// again, because a second Cancel would send the scope TERM again. A +// program that Cancel never reached keeps running until it exits or the +// Session ends. +// - A broker lost after the acknowledgement makes the shim exit with 255, +// with the reason on its stderr when the relay can still write it. A +// relay lost after the acknowledgement makes the shim exit with 255 and +// no message, because the shim no longer holds its stderr. +// - A pipe, FIFO or pty slave must be open for the direction the program +// uses it in; the shim fails with 126 otherwise. Reads and writes of a +// regular file or block device block, as a native program's do. A +// character device other than a pty slave, such as /dev/tty, and a pipe, +// FIFO or pty slave that the Session user may not reopen, is shared: the +// relay polls it before each call and writes at most PIPE_BUF bytes at +// once, and a read still waits when another reader took the data the +// poll reported. +// - Output on an AF_UNIX socket names the relay's pid, uid and gid, not +// the shim's. +// - A signal sent to the shim reaches the remote program only when the +// process service declares it. The shim catches every signal a Go +// program can catch except CHLD, PIPE, URG and PROF, and the broker drops +// the ones the service does not declare. Of the signals it does not +// catch, KILL ends the shim, ILL, TRAP, BUS, FPE, SEGV, STKFLT and SYS +// make the Go runtime end it with status 2, and signals 32 and 34 end it; +// a shim ended before the exit cancels the operation. PIPE, PROF and +// signal 33 have no effect. +// - A signal ignored when the shim started is still forwarded unless it is +// HUP or INT, because the Go runtime replaces inherited ignores. +// - A signal sent to the shim's PID reaches the remote initial process +// group, or on a terminal the foreground process group for INT, QUIT, +// TSTP, TTIN, TTOU, CONT and HUP, because the shim cannot tell it from a +// signal sent to its process group. +package processbroker diff --git a/apps/daemon/internal/processbroker/fakeservice_linux_test.go b/apps/daemon/internal/processbroker/fakeservice_linux_test.go new file mode 100644 index 00000000..0e1e1a95 --- /dev/null +++ b/apps/daemon/internal/processbroker/fakeservice_linux_test.go @@ -0,0 +1,410 @@ +//go:build linux + +package processbroker + +import ( + "context" + "io" + "net" + "slices" + "sync" + + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// fakeService is a process service whose program copies stdin to stdout and +// exits 0 when stdin closes. Its knobs reproduce one service behavior each. +type fakeService struct { + instance sandboxwire.ID + caps sp.Capabilities + + // startLate keeps an operation Starting until runStarting, or until + // stdin is written, which it refuses. attachedStarting closes once an + // Attach finds the operation Starting. + startLate bool + attachedStarting chan struct{} + // startFails ends the start with StartFailed instead of Started. + startFails bool + // partialFirst accepts half of the first stdin write. + partialFirst bool + // inspectFails refuses Inspect for good. + inspectFails bool + // leaderExits exits the leader with 0 as it takes the first stdin + // write. A background process it leaves copies the rest of stdin to + // stdout until end of file. + leaderExits bool + + mu sync.Mutex + ops map[sandboxwire.ID]*fakeOp + writes int // stdin writes the service took + cancels int + refusedIn int // stdin requests refused while Starting +} + +type fakeOp struct { + id sandboxwire.ID + state sp.OperationState + events []sp.Event // events[i] has sequence i+1 + changed chan struct{} + stdin []byte + stdinClosed bool + closedAt uint64 // the CloseStdin offset + stdout uint64 // the stdout offset + exit *sp.ExitStatus + failure *sp.Failure + background bool // a background process outlives the leader + released bool +} + +func newFakeService() *fakeService { + return &fakeService{ + instance: sandboxwire.NewID(), + caps: sp.Capabilities{ + Platform: sp.PlatformLinux, + Scopes: []sp.Scope{sp.ScopePOSIXSession}, + IOModes: []sp.IOMode{sp.IOPipes}, + Signals: []sp.Signal{1, 2, 15}, + SignalTargets: []sp.SignalTarget{sp.TargetInitialProcessGroup}, + MaxStartBytes: sandboxwire.MaxPayload, + MaxDataBytes: sandboxwire.MaxChunk, + MaxActiveOperations: 8, + MaxOperationRecords: 8, + MaxReplayBytesPerOperation: 1 << 20, + OwnerLossGraceMillis: 60000, + CancelGraceLimitMillis: 60000, + }, + ops: map[sandboxwire.ID]*fakeOp{}, + attachedStarting: make(chan struct{}), + } +} + +// serve replaces each stream to the real service with one to s, and loses +// the first stream in place of the response to the first request that lose +// selects. +func (s *fakeService) serve(ctx context.Context, lose func(sandboxwire.Frame) bool) func(int32, net.Conn) io.ReadWriteCloser { + return func(n int32, c net.Conn) io.ReadWriteCloser { + c.Close() + broker, svc := net.Pipe() + go sp.Serve(ctx, svc, sp.Attachment{ID: s.instance}, s) + if n > 1 || lose == nil { + return broker + } + return intercept(broker, func(fr sandboxwire.Frame) verdict { + if lose(fr) { + return loseResponse + } + return pass + }) + } +} + +// emit appends an event; s.mu is held. +func (s *fakeService) emit(op *fakeOp, ev func(sp.EventHeader) sp.Event) { + op.events = append(op.events, ev(sp.EventHeader{OperationID: op.id, Sequence: uint64(len(op.events) + 1)})) + close(op.changed) + op.changed = make(chan struct{}) +} + +// subscribe sends op's events after seq on conn until the stream ends. +func (s *fakeService) subscribe(conn *sp.Conn, op *fakeOp, after uint64) { + go func() { + next := after + for { + s.mu.Lock() + evs := slices.Clone(op.events[min(next, uint64(len(op.events))):]) + changed := op.changed + s.mu.Unlock() + for _, ev := range evs { + if conn.Send(ev) != nil { + return + } + next++ + } + select { + case <-changed: + case <-conn.Context().Done(): + return + } + } + }() +} + +// run ends the start; s.mu is held. +func (s *fakeService) run(op *fakeOp) { + if op.state != sp.StateStarting { + return + } + if s.startFails { + op.state = sp.StateStartFailed + op.failure = sp.Fail(sp.CodeNotFound, sandboxwire.EffectNone, "no such file") + f := *op.failure + s.emit(op, func(h sp.EventHeader) sp.Event { return sp.StartFailedEvent{EventHeader: h, Failure: f} }) + return + } + op.state = sp.StateRunning + s.emit(op, func(h sp.EventHeader) sp.Event { return sp.StartedEvent{EventHeader: h} }) +} + +// runStarting ends every start in progress. +func (s *fakeService) runStarting() { + s.mu.Lock() + defer s.mu.Unlock() + for _, op := range s.ops { + s.run(op) + } +} + +// exited ends the program with status, or only its background process +// when the leader already exited; s.mu is held. +func (s *fakeService) exited(op *fakeOp, status sp.ExitStatus) { + switch { + case op.background: + op.background = false + s.drained(op, nil) + case op.state == sp.StateRunning: + op.state, op.exit = sp.StateExited, &status + s.drained(op, &status) + } +} + +// drained closes the output and the scope, with the leader's exit between +// them when it exits now; s.mu is held. +func (s *fakeService) drained(op *fakeOp, exit *sp.ExitStatus) { + out := op.stdout + s.emit(op, func(h sp.EventHeader) sp.Event { + return sp.StreamClosedEvent{EventHeader: h, Stream: sp.StreamStdout, Offset: out, Disposition: sp.OutputDrained} + }) + s.emit(op, func(h sp.EventHeader) sp.Event { + return sp.StreamClosedEvent{EventHeader: h, Stream: sp.StreamStderr, Disposition: sp.OutputDrained} + }) + if exit != nil { + s.emit(op, func(h sp.EventHeader) sp.Event { return sp.ExitedEvent{EventHeader: h, Status: *exit} }) + } + s.emit(op, func(h sp.EventHeader) sp.Event { + return sp.OutputClosedEvent{EventHeader: h, Disposition: sp.OutputDrained} + }) + s.emit(op, func(h sp.EventHeader) sp.Event { return sp.ScopeClosedEvent{EventHeader: h} }) +} + +func (s *fakeService) status(op *fakeOp) sp.OperationStatus { + st := sp.OperationStatus{ + State: op.state, Exit: op.exit, StartFailure: op.failure, + StdinOffset: uint64(len(op.stdin)), StdinClosed: op.stdinClosed, + Scope: sp.ScopeStateActive, Released: op.released, + FirstRetained: 1, LastSequence: uint64(len(op.events)), + } + if op.state == sp.StateExited && !op.background || op.state == sp.StateStartFailed { + d := sp.OutputDrained + st.Output, st.Scope = &d, sp.ScopeStateClosed + } + return st +} + +// lookup returns the operation with s.mu held. +func (s *fakeService) lookup(ref sp.OperationRef) (*fakeOp, error) { + if ref.ServerInstanceID != s.instance { + return nil, sp.Fail(sp.CodeInstanceChanged, sandboxwire.EffectNone, "instance changed") + } + s.mu.Lock() + op := s.ops[ref.OperationID] + if op == nil { + s.mu.Unlock() + return nil, sp.Fail(sp.CodeNotFound, sandboxwire.EffectNone, "no operation") + } + if op.released { + s.mu.Unlock() + return nil, sp.Fail(sp.CodeReleased, sandboxwire.EffectNone, "released") + } + return op, nil +} + +func (s *fakeService) Describe(context.Context, *sp.Conn, sp.DescribeRequest) (sp.DescribeResponse, error) { + return sp.DescribeResponse{ServerInstanceID: s.instance, Capabilities: s.caps}, nil +} + +func (s *fakeService) Start(_ context.Context, conn *sp.Conn, req sp.StartRequest) (sp.StartResponse, error) { + if req.ServerInstanceID != s.instance { + return sp.StartResponse{}, sp.Fail(sp.CodeInstanceChanged, sandboxwire.EffectNone, "instance changed") + } + s.mu.Lock() + defer s.mu.Unlock() + if s.ops[req.OperationID] != nil { + return sp.StartResponse{Disposition: sp.StartExisting}, nil + } + op := &fakeOp{id: req.OperationID, state: sp.StateStarting, changed: make(chan struct{})} + s.ops[op.id] = op + s.subscribe(conn, op, 0) + if !s.startLate { + s.run(op) + } + return sp.StartResponse{Disposition: sp.StartCreated}, nil +} + +func (s *fakeService) Attach(_ context.Context, conn *sp.Conn, req sp.AttachRequest) (sp.AttachResponse, error) { + op, err := s.lookup(req.OperationRef) + if err != nil { + return sp.AttachResponse{}, err + } + defer s.mu.Unlock() + s.subscribe(conn, op, req.AfterSequence) + if op.state == sp.StateStarting { + select { + case <-s.attachedStarting: + default: + close(s.attachedStarting) + } + } + return sp.AttachResponse{Status: s.status(op)}, nil +} + +func (s *fakeService) Inspect(_ context.Context, _ *sp.Conn, req sp.InspectRequest) (sp.InspectResponse, error) { + op, err := s.lookup(req.OperationRef) + if err != nil { + return sp.InspectResponse{}, err + } + defer s.mu.Unlock() + if s.inspectFails { + return sp.InspectResponse{}, sp.Fail(sp.CodeIO, sandboxwire.EffectNone, "inspect failed") + } + return sp.InspectResponse{Status: s.status(op)}, nil +} + +// stdinRefused refuses stdin while the operation starts, which also ends a +// late start; s.mu is held. +func (s *fakeService) stdinRefused(op *fakeOp) error { + switch { + case op.state == sp.StateStarting: + s.refusedIn++ + s.run(op) + return sp.Fail(sp.CodeNotRunning, sandboxwire.EffectNone, "the operation is starting") + case op.state != sp.StateRunning && !op.background: + return sp.Fail(sp.CodeNotRunning, sandboxwire.EffectNone, "the operation is not running") + case op.stdinClosed: + return sp.Fail(sp.CodeStdinClosed, sandboxwire.EffectNone, "stdin is closed") + } + return nil +} + +func (s *fakeService) WriteStdin(_ context.Context, _ *sp.Conn, req sp.WriteStdinRequest) (sp.WriteStdinResponse, error) { + op, err := s.lookup(req.OperationRef) + if err != nil { + return sp.WriteStdinResponse{}, err + } + defer s.mu.Unlock() + if err := s.stdinRefused(op); err != nil { + return sp.WriteStdinResponse{}, err + } + if req.Offset != uint64(len(op.stdin)) { + return sp.WriteStdinResponse{}, sp.Fail(sp.CodeInputOffsetConflict, sandboxwire.EffectNone, "offset %d is not %d", req.Offset, len(op.stdin)) + } + data := req.Data + if s.writes++; s.writes == 1 && s.partialFirst { + data = data[:len(data)/2] + } + op.stdin = append(op.stdin, data...) + if len(data) > 0 { + out, echo := op.stdout, slices.Clone(data) + op.stdout += uint64(len(data)) + s.emit(op, func(h sp.EventHeader) sp.Event { + return sp.OutputEvent{EventHeader: h, Stream: sp.StreamStdout, Offset: out, Data: echo} + }) + } + if s.leaderExits && op.state == sp.StateRunning { + status := sp.ExitStatus{Kind: sp.ExitCode} + op.state, op.exit, op.background = sp.StateExited, &status, true + s.emit(op, func(h sp.EventHeader) sp.Event { return sp.ExitedEvent{EventHeader: h, Status: status} }) + } + return sp.WriteStdinResponse{Accepted: uint32(len(data))}, nil +} + +func (s *fakeService) CloseStdin(_ context.Context, _ *sp.Conn, req sp.CloseStdinRequest) (sp.CloseStdinResponse, error) { + op, err := s.lookup(req.OperationRef) + if err != nil { + return sp.CloseStdinResponse{}, err + } + defer s.mu.Unlock() + if op.stdinClosed && req.Offset == op.closedAt { + return sp.CloseStdinResponse{}, nil + } + if err := s.stdinRefused(op); err != nil { + return sp.CloseStdinResponse{}, err + } + if req.Offset != uint64(len(op.stdin)) { + return sp.CloseStdinResponse{}, sp.Fail(sp.CodeInputOffsetConflict, sandboxwire.EffectNone, "offset %d is not %d", req.Offset, len(op.stdin)) + } + op.stdinClosed, op.closedAt = true, req.Offset + s.exited(op, sp.ExitStatus{Kind: sp.ExitCode}) + return sp.CloseStdinResponse{}, nil +} + +func (s *fakeService) CloseOutput(_ context.Context, _ *sp.Conn, req sp.CloseOutputRequest) (sp.CloseOutputResponse, error) { + if _, err := s.lookup(req.OperationRef); err != nil { + return sp.CloseOutputResponse{}, err + } + s.mu.Unlock() + return sp.CloseOutputResponse{}, nil +} + +func (s *fakeService) ResizePTY(context.Context, *sp.Conn, sp.ResizePTYRequest) (sp.ResizePTYResponse, error) { + return sp.ResizePTYResponse{}, sp.Fail(sp.CodeUnsupported, sandboxwire.EffectNone, "no terminal") +} + +func (s *fakeService) Signal(_ context.Context, _ *sp.Conn, req sp.SignalRequest) (sp.SignalResponse, error) { + if _, err := s.lookup(req.OperationRef); err != nil { + return sp.SignalResponse{}, err + } + s.mu.Unlock() + return sp.SignalResponse{}, nil +} + +// Cancel ends the program as TERM does. +func (s *fakeService) Cancel(_ context.Context, _ *sp.Conn, req sp.CancelRequest) (sp.CancelResponse, error) { + op, err := s.lookup(req.OperationRef) + if err != nil { + return sp.CancelResponse{}, err + } + defer s.mu.Unlock() + s.cancels++ + s.exited(op, sp.ExitStatus{Kind: sp.ExitSignal, Signal: 15}) + return sp.CancelResponse{}, nil +} + +func (s *fakeService) AckEvents(_ context.Context, _ *sp.Conn, req sp.AckEventsRequest) (sp.AckEventsResponse, error) { + if _, err := s.lookup(req.OperationRef); err != nil { + return sp.AckEventsResponse{}, err + } + s.mu.Unlock() + return sp.AckEventsResponse{}, nil +} + +func (s *fakeService) Release(_ context.Context, _ *sp.Conn, req sp.ReleaseRequest) (sp.ReleaseResponse, error) { + op, err := s.lookup(req.OperationRef) + if err != nil { + return sp.ReleaseResponse{}, err + } + defer s.mu.Unlock() + if op.state != sp.StateExited && op.state != sp.StateStartFailed || op.background { + return sp.ReleaseResponse{}, sp.Fail(sp.CodeBusy, sandboxwire.EffectNone, "not settled") + } + op.released = true + return sp.ReleaseResponse{}, nil +} + +// fakeCounts is what a fakeService saw. +type fakeCounts struct { + writes, cancels, refusedIn int + stdin []byte // the only operation's accepted stdin + closedAt uint64 // and its CloseStdin offset + released bool // and whether it was released +} + +func (s *fakeService) counts() fakeCounts { + s.mu.Lock() + defer s.mu.Unlock() + c := fakeCounts{writes: s.writes, cancels: s.cancels, refusedIn: s.refusedIn} + for _, op := range s.ops { + c.stdin, c.closedAt, c.released = slices.Clone(op.stdin), op.closedAt, op.released + } + return c +} diff --git a/apps/daemon/internal/processbroker/invocation_linux.go b/apps/daemon/internal/processbroker/invocation_linux.go new file mode 100644 index 00000000..eea43903 --- /dev/null +++ b/apps/daemon/internal/processbroker/invocation_linux.go @@ -0,0 +1,874 @@ +//go:build linux + +package processbroker + +import ( + "context" + "errors" + "fmt" + "log/slog" + "path" + "slices" + "sync" + "time" + + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// invocation is one shim invocation: the operation it started and the +// streams the broker forwards for it through the relay. +type invocation struct { + b *Broker + rid uint64 // the relay's invocation ID + open processshim.Open + log *slog.Logger + spec sp.ProcessSpec + id sandboxwire.ID + term *processshim.Terminal // nil for pipes + + // halt ends every wait of the invocation; stopIn ends stdin forwarding. + halt chan struct{} + haltOnce sync.Once + stopIn chan struct{} + stopOnce sync.Once + + writing sync.WaitGroup // the output writers + helpers sync.WaitGroup // everything else but the stdin pump + started chan struct{} // closed once the operation exists or never will + // gone ends when the shim is lost. + gone context.Context + loseShim context.CancelFunc + + acks tracker + writers map[sp.Stream]*writer + byFD [3]*writer + sigs chan processshim.Signaled + // input holds the relay's answer to the outstanding Read. + input chan processshim.RelayMessage + // running is set once Started arrives; only observe uses it. + running bool + + // sendMu orders the invocation's messages before its End. + sendMu sync.Mutex + ended bool + + mu sync.Mutex + inst sandboxwire.ID + cur handle + exited bool // the exit is decided: Exited, StartFailed or exit lost + shimLost bool + replied bool + credit uint32 // the outstanding Read's Max; 0 for none + // settlement, from delivered events + startFailed, outputClosed, scopeClosed bool +} + +// handle is the operation's handle on one stream. relinked closes when a +// re-Attach replaces it. +type handle struct { + op *sp.Operation + s *stream + relinked chan struct{} +} + +// ptyGroupSignals target the terminal's foreground group on a PTY. +var ptyGroupSignals = []uint16{ + uint16(unix.SIGINT), uint16(unix.SIGQUIT), uint16(unix.SIGTSTP), uint16(unix.SIGTTIN), + uint16(unix.SIGTTOU), uint16(unix.SIGCONT), uint16(unix.SIGHUP), +} + +func (b *Broker) newInvocation(open processshim.Open) *invocation { + inv := &invocation{ + b: b, rid: open.ID, open: open, log: b.log.With("invocation", open.ID), + id: sandboxwire.NewID(), term: open.Terminal, + halt: make(chan struct{}), stopIn: make(chan struct{}), started: make(chan struct{}), + sigs: make(chan processshim.Signaled, 64), input: make(chan processshim.RelayMessage, 1), + writers: map[sp.Stream]*writer{}, + } + inv.gone, inv.loseShim = context.WithCancel(context.Background()) + inv.acks.init() + add := func(stream sp.Stream, fd uint8) { + w := &writer{inv: inv, stream: stream, fd: fd, wake: make(chan struct{}, 1)} + inv.writers[stream] = w + inv.byFD[fd] = w + } + if inv.term != nil { + add(sp.StreamTerminal, 1) + } else { + add(sp.StreamStdout, 1) + add(sp.StreamStderr, 2) + } + return inv +} + +// serve refuses the invocation or accepts and runs it, then ends it. +func (inv *invocation) serve() { + defer inv.b.wg.Done() + defer inv.teardown() + if refusal := inv.prepare(); refusal != nil { + inv.reply(*refusal, nil) + return + } + stop := context.AfterFunc(inv.b.ctx, inv.halted) + defer stop() + if inv.send(processshim.Accept{ID: inv.rid}) != nil { + return + } + inv.run() +} + +func refuse(code uint8, format string, args ...any) *processshim.Result { + return &processshim.Result{Code: code, Message: message(fmt.Sprintf(format, args...))} +} + +// message is msg within the IPC's limit. +func message(msg string) []byte { + if msg == "" { + msg = "failed" + } + return []byte(msg[:min(len(msg), processshim.MaxMessageBytes)]) +} + +// prepare checks the request and builds the spec. Every refusal happens +// here, before the relay acknowledges the shim. +func (inv *invocation) prepare() *processshim.Result { + b, req := inv.b, inv.open.Request + remote, ok := b.cfg.Executables.resolve(string(req.ExecPath), string(req.Cwd)) + if !ok { + return refuse(processshim.ExitNotFound, "%s: not a declared sandbox executable", req.ExecPath) + } + if underPrivate(path.Clean(string(req.Cwd))) { + return refuse(processshim.ExitCannotRun, "%s: the working directory is private to the Session", req.Cwd) + } + inv.log = inv.log.With("executable", remote) + env, dropped := b.cfg.Environment.compose(req.Env) + if len(dropped) > 0 { + inv.log.Info("environment entries naming the private directory dropped", "names", dropped) + } + spec := sp.ProcessSpec{ + Executable: []byte(remote), + Argv: req.Argv, + Env: env, + Cwd: req.Cwd, + Umask: req.Umask, + IOMode: sp.IOPipes, + Scope: b.cfg.Scope, + } + if t := inv.term; t != nil { + spec.IOMode = sp.IOPTY + spec.PTY = &sp.PTYSpec{Size: windowSize(t.Size), Term: termName(spec.Env)} + spec.Env = slices.DeleteFunc(spec.Env, func(v sp.EnvVar) bool { return string(v.Name) == "TERM" }) + } + if err := spec.Validate(); err != nil { + return refuse(processshim.ExitCannotRun, "%s: %v", remote, err) + } + inv.spec = spec + return nil +} + +func windowSize(s processshim.WindowSize) sp.WindowSize { + return sp.WindowSize{Rows: s.Rows, Cols: s.Cols, XPixels: s.XPixels, YPixels: s.YPixels} +} + +// termios is the terminal's saved mode. +func termios(t *processshim.Terminal) *unix.Termios { + tio := &unix.Termios{Iflag: t.Iflag, Oflag: t.Oflag, Cflag: t.Cflag, Lflag: t.Lflag} + copy(tio.Cc[:], t.Cc) + return tio +} + +// termName is the remote TERM: the composed environment's, unless it is +// unusable. +func termName(env []sp.EnvVar) []byte { + for _, v := range env { + if string(v.Name) == "TERM" && len(v.Value) > 0 && len(v.Value) <= sp.MaxTermBytes { + return v.Value + } + } + return []byte("dumb") +} + +// run observes the started operation. The program runs, and stdin is +// forwarded, only once its Started event arrives: an operation that a +// retried Start found may still be starting, and refuses stdin until then. +func (inv *invocation) run() { + h, ok := inv.start() + close(inv.started) + if !ok { + return + } + inv.writing.Add(len(inv.writers)) + for _, w := range inv.writers { + go w.run() + } + inv.helpers.Add(2) + go inv.ackLoop() + go inv.forwardSignals() + inv.observe(h) +} + +// resolveWindow bounds how long a Start that may have taken effect is +// resolved before the shim gets 255. +const resolveWindow = 30 * time.Second + +// startNext says what follows one Start attempt. +type startNext int + +const ( + startDone startNext = iota // the operation is observed + startEnded // the invocation ended + startNow // start again on the next stream + startLater // start again after a backoff +) + +// start starts the operation. A Start that may have taken effect is retried +// with the same ID and spec until its outcome is definite: the service then +// answers with the existing operation, or refuses in a way that proves none +// exists. One deadline, set by the first Start that may have taken effect, +// bounds every later link wait, Start, implicit Attach and backoff; when it +// passes, the shim gets 255. A Start that had no effect may move to a new +// service incarnation. Until a Start may have taken effect, losing the shim +// ends the invocation, including a wait for a stream. +func (inv *invocation) start() (handle, bool) { + var deadline time.Time // set once a Start may have taken effect + exists := false // a Start found the operation + backoff := minBackoff + for { + if !deadline.IsZero() && !time.Now().Before(deadline) { + inv.unconfirmed() + return handle{}, false + } + h, next := inv.startOnce(&deadline, &exists) + switch next { + case startDone: + return h, true + case startEnded: + return handle{}, false + case startNow: + continue + } + wait, gone := backoff, inv.gone.Done() + if !deadline.IsZero() { + wait, gone = min(wait, time.Until(deadline)), nil + } + if !inv.sleepUnless(wait, gone) { + inv.fail("the process broker stopped") + return handle{}, false + } + backoff = min(2*backoff, maxBackoff) + } +} + +// startOnce makes one Start attempt. A deadline it sets or finds bounds the +// attempt, including the Start's implicit Attach. Until a Start may have +// taken effect, the shim's loss ends the wait for a stream. +func (inv *invocation) startOnce(deadline *time.Time, exists *bool) (handle, startNext) { + var ctx context.Context + var cancel context.CancelFunc + if deadline.IsZero() { + ctx, cancel = context.WithCancel(inv.b.ctx) + defer context.AfterFunc(inv.gone, cancel)() + } else { + ctx, cancel = context.WithDeadline(inv.b.ctx, *deadline) + } + defer cancel() + s, err := inv.b.link.get(ctx) + switch { + case inv.b.ctx.Err() != nil: + inv.fail("the process broker stopped") + return handle{}, startEnded + case deadline.IsZero() && inv.lost(): + return handle{}, startEnded // nothing started + case err != nil: + return handle{}, startLater // the deadline passed + } + if inv.inst != s.instance { + if !deadline.IsZero() { + inv.fail("the sandbox process service restarted while the program was starting") + return handle{}, startEnded + } + inv.inst = s.instance + } + if inv.spec.PTY != nil && inv.spec.PTY.Modes == nil { + inv.spec.PTY.Modes = sp.ReadModes(termios(inv.term), s.caps.PTYModes) + } + if f := s.caps.CheckStart(inv.spec); f != nil { + inv.reply(*refuse(processshim.ExitCannotRun, "%s: %s", inv.spec.Executable, f.Message), nil) + return handle{}, startEnded + } + req := sp.StartRequest{OperationRef: sp.OperationRef{ServerInstanceID: inv.inst, OperationID: inv.id}, Spec: inv.spec} + if n := len(sp.Encode(req)); n > int(s.caps.MaxStartBytes) { + inv.reply(*refuse(processshim.ExitCannotRun, "%s: argument list and environment of %d bytes exceed %d", inv.spec.Executable, n, s.caps.MaxStartBytes), nil) + return handle{}, startEnded + } + began := time.Now() + if deadline.IsZero() { + var cancelStart context.CancelFunc + ctx, cancelStart = context.WithDeadline(inv.b.ctx, began.Add(resolveWindow)) + defer cancelStart() + } + op, disp, err := s.client.Start(ctx, inv.inst, inv.id, inv.spec) + if err == nil { + h := handle{op: op, s: s, relinked: make(chan struct{})} + inv.mu.Lock() + inv.cur = h + inv.mu.Unlock() + return h, startDone + } + f := asFailure(err) + if inv.b.ctx.Err() != nil { + inv.fail("the process broker stopped") + return handle{}, startEnded + } + if disp == sp.StartExisting { + *exists = true // its implicit Attach failed; the next Start attaches again + } + if deadline.IsZero() && (*exists || f.Effect == sandboxwire.EffectPossible) { + *deadline = began.Add(resolveWindow) + } + if deadline.IsZero() { // nothing has started + switch { + case s.ended(): + return handle{}, startNow + case f.Code == sp.CodeBusy: + return handle{}, startLater + } + inv.reply(*refuse(processshim.ExitCannotRun, "%s: %s", inv.spec.Executable, f.Message), nil) + return handle{}, startEnded + } + switch { + case f.Code == sp.CodeInstanceChanged: + inv.fail("the sandbox process service restarted while the program was starting") + return handle{}, startEnded + case f.Code == sp.CodeReleased || f.Code == sp.CodeOperationConflict: + inv.fail(fmt.Sprintf("the program's start could not be resolved: %s", f.Message)) + return handle{}, startEnded + case !*exists && f.Effect == sandboxwire.EffectNone && provesAbsence(f.Code): + inv.reply(*refuse(processshim.ExitCannotRun, "%s: %s", inv.spec.Executable, f.Message), nil) + return handle{}, startEnded + } + return handle{}, startLater +} + +// provesAbsence reports whether a Start refusal shows that the ID has no +// operation. A service answers a Start of an existing ID and spec with +// Existing, so a refusal its Start handling makes after that lookup proves +// absence. Busy, InvalidArgument from decoding, and the failures of the +// client and the stream come before the lookup and prove nothing. +func provesAbsence(c sp.ErrorCode) bool { + switch c { + case sp.CodeStaleAttachment, sp.CodeResourceExhausted, sp.CodeUnsupported: + return true + } + return false +} + +// sleep waits for d and reports false when the invocation halts first. +func (inv *invocation) sleep(d time.Duration) bool { return inv.sleepUnless(d, nil) } + +// sleepUnless is sleep that also ends, reporting true, when wake closes. +func (inv *invocation) sleepUnless(d time.Duration, wake <-chan struct{}) bool { + t := time.NewTimer(d) + defer t.Stop() + select { + case <-t.C: + return true + case <-wake: + return true + case <-inv.halt: + return false + } +} + +// unconfirmed ends an invocation whose start never became definite, and +// cancels the operation if it exists after all. A Start that takes effect +// later is left to the Session's cleanup. +func (inv *invocation) unconfirmed() { + inv.fail("the program may have started, but its start could not be confirmed") + s := inv.b.link.current() + if s == nil || s.instance != inv.inst { + return + } + ctx, cancel := context.WithTimeout(inv.b.ctx, 5*time.Second) + defer cancel() + op, _, err := s.client.Attach(ctx, inv.inst, inv.id, 0) + if err != nil { + return + } + if err := op.Cancel(ctx, uint32(inv.b.cfg.CancelGrace.Milliseconds())); err != nil { + inv.log.Warn("operation cancel failed", "error", err) + } + op.Detach() +} + +// observe handles events until the operation is released, then returns. +func (inv *invocation) observe(h handle) { + var received uint64 + for { + select { + case ev, ok := <-h.op.Events(): + if !ok { + next, done := inv.reattach(received) + if done { + return + } + h = next + // A settled operation has no events left to resume. + if inv.isSettled() && inv.settle(h, received) { + return + } + continue + } + received = ev.Header().Sequence + if inv.handle(h, ev) { + return + } + case <-inv.halt: + inv.fail("the process broker stopped") + h.op.Detach() + return + } + } +} + +// reattach resumes the operation after the last received event on a new +// stream. It reports done when the operation is over for the broker. +func (inv *invocation) reattach(after uint64) (handle, bool) { + backoff := minBackoff + for { + s, err := inv.b.link.get(inv.b.ctx) + if err != nil { + inv.fail("the process broker stopped") + return handle{}, true + } + if s.instance != inv.inst { + inv.fail("the sandbox process service restarted; the program's outcome is unknown") + return handle{}, true + } + op, st, err := s.client.Attach(inv.b.ctx, inv.inst, inv.id, after) + if err != nil { + if s.ended() { + continue + } + f := asFailure(err) + switch { + case refused(f): + if inv.sleep(backoff) { + backoff = min(2*backoff, maxBackoff) + continue + } + inv.fail("the process broker stopped") + case f.Code == sp.CodeReleased: + return handle{}, true + case f.Code == sp.CodeReplayGap: + inv.fail("output was lost while the sandbox was unreachable") + default: + inv.fail(fmt.Sprintf("the operation could not be resumed: %s", f.Message)) + } + return handle{}, true + } + if st.Released { + op.Detach() + return handle{}, true + } + inv.mu.Lock() + old := inv.cur + inv.cur = handle{op: op, s: s, relinked: make(chan struct{})} + h := inv.cur + inv.mu.Unlock() + close(old.relinked) + return h, false + } +} + +// handle applies one event and reports whether the operation is done. +func (inv *invocation) handle(h handle, ev sp.Event) bool { + seq := ev.Header().Sequence + switch ev := ev.(type) { + case sp.StartedEvent: + if !inv.running { + inv.running = true + inv.send(processshim.Started{ID: inv.rid}) + go inv.pumpStdin(h.s.caps) + } + case sp.OutputEvent: + if w := inv.writers[ev.Stream]; w != nil { + w.push(chunk{seq: seq, data: ev.Data}) + return false + } + case sp.StreamClosedEvent: + if w := inv.writers[ev.Stream]; w != nil { + w.push(chunk{seq: seq, close: true}) + return false + } + case sp.ExitedEvent: + inv.decideExit() + // The relay answers the shim once the output the program wrote + // before exiting is written, as a native exit follows its writes. + inv.reply(exitResult(ev.Status), inv.marks()) + case sp.StartFailedEvent: + inv.decideExit() + code := uint8(processshim.ExitCannotRun) + if ev.Failure.Code == sp.CodeNotFound { + code = processshim.ExitNotFound + } + inv.reply(*refuse(code, "%s: %s", inv.spec.Executable, ev.Failure.Message), nil) + inv.setSettlement(func() { inv.startFailed = true }) + case sp.ObservationLostEvent: + if ev.Observation == sp.ObservationExit { + inv.decideExit() + inv.reply(*refuse(processshim.ExitLost, "the program's exit status was lost: %s", ev.Failure.Message), nil) + } else { + // The service keeps watching the scope and reports its close. + inv.log.Info("operation scope observation lost", "reason", ev.Failure.Message) + } + case sp.OutputClosedEvent: + inv.setSettlement(func() { inv.outputClosed = true }) + case sp.ScopeClosedEvent: + inv.setSettlement(func() { inv.scopeClosed = true }) + } + inv.acks.deliver(seq) + if inv.isSettled() { + return inv.settle(h, seq) + } + return false +} + +func exitResult(s sp.ExitStatus) processshim.Result { + if s.Kind == sp.ExitSignal { + return processshim.Result{Signal: uint16(s.Signal)} + } + return processshim.Result{Code: s.Code} +} + +func (inv *invocation) setSettlement(set func()) { + inv.mu.Lock() + set() + inv.mu.Unlock() +} + +// isSettled mirrors the service's settlement: Release succeeds once it holds. +func (inv *invocation) isSettled() bool { + inv.mu.Lock() + defer inv.mu.Unlock() + return inv.startFailed || (inv.exited && inv.outputClosed && inv.scopeClosed) +} + +// settle releases the settled operation once every event through seq +// reached its destination, so Release follows the delivery of Exited and +// OutputClosed. It reports false when the stream ended first; the caller +// re-Attaches and settles again. +func (inv *invocation) settle(h handle, seq uint64) bool { + if !inv.acks.wait(seq, inv.halt) { + h.op.Detach() + return true + } + return inv.releaseOp(h) +} + +// releaseOp releases the settled operation. A Busy refusal is retried +// until the invocation halts. It reports false when the stream ended first. +func (inv *invocation) releaseOp(h handle) bool { + for backoff := minBackoff; ; backoff = min(2*backoff, maxBackoff) { + err := h.op.Release(inv.b.ctx) + if err == nil { + return true + } + if h.s.ended() && inv.b.ctx.Err() == nil { + return false + } + f := asFailure(err) + if f.Code == sp.CodeBusy && inv.sleep(backoff) { + continue + } + if f.Code != sp.CodeReleased && f.Code != sp.CodeBusy { + inv.log.Warn("operation release failed", "error", err) + } + h.op.Detach() + return true + } +} + +func (inv *invocation) decideExit() { + inv.mu.Lock() + inv.exited = true + inv.mu.Unlock() + inv.stopInput() +} + +func (inv *invocation) lost() bool { + inv.mu.Lock() + defer inv.mu.Unlock() + return inv.shimLost +} + +func (inv *invocation) current() handle { + inv.mu.Lock() + defer inv.mu.Unlock() + return inv.cur +} + +// relink waits for the handle that replaces h after its stream ended. +func (inv *invocation) relink(h handle) (handle, bool) { + select { + case <-h.relinked: + return inv.current(), true + case <-inv.halt: + return h, false + } +} + +// reply has the relay send the shim its one Result once the marked output +// is written, unless the shim has a Result or is gone. It reports false +// then. A Result Message reaches the shim's stderr. +func (inv *invocation) reply(r processshim.Result, marks []processshim.Mark) bool { + return inv.answer(r, marks, false) +} + +// replyBeforeExit is reply that also reports false once the exit is decided. +// It checks the exit and reserves the Result in one step, so the Result +// never replaces a decided exit's. +func (inv *invocation) replyBeforeExit(r processshim.Result) bool { + return inv.answer(r, nil, true) +} + +func (inv *invocation) answer(r processshim.Result, marks []processshim.Mark, beforeExit bool) bool { + inv.mu.Lock() + if inv.replied || inv.shimLost || beforeExit && inv.exited { + inv.mu.Unlock() + return false + } + inv.replied = true + inv.mu.Unlock() + inv.stopInput() + inv.send(processshim.Exit{ID: inv.rid, Result: r, Marks: marks}) + return true +} + +// fail ends the invocation with 255 and the reason on the shim's stderr. The +// reason is written even after the shim exited, because the Harness may still +// read the output the failure cut short. +func (inv *invocation) fail(reason string) { + inv.log.Warn("process invocation failed", "reason", reason) + inv.halted() + if !inv.reply(processshim.Result{Code: processshim.ExitLost, Message: message(reason)}, nil) { + inv.send(processshim.Notice{ID: inv.rid, Message: message(reason)}) + } +} + +// shimGone handles the end of the shim's connection. Before the exit is +// decided it cancels the operation; afterwards nothing is cancelled, as a +// native background job survives its parent. +func (inv *invocation) shimGone() { + inv.mu.Lock() + if inv.replied || inv.shimLost { + inv.mu.Unlock() + return + } + inv.shimLost = true + inv.loseShim() + cancel := !inv.exited + inv.mu.Unlock() + inv.log.Info("process shim lost", "cancel", cancel) + inv.stopInput() + if cancel { + inv.cancelRemote() + } +} + +func (inv *invocation) cancelRemote() { + select { + case <-inv.started: + case <-inv.halt: + return + } + h := inv.current() + if h.op == nil { + return + } + grace := uint32(inv.b.cfg.CancelGrace.Milliseconds()) + err := inv.request(h, false, func(h handle) error { + if inv.exitDecided() { + return nil + } + return h.op.Cancel(inv.b.ctx, grace) + }) + switch { + case err == nil || inv.b.ctx.Err() != nil: + case asFailure(err).Effect == sandboxwire.EffectPossible: + // A second Cancel would send the scope TERM again; the events tell + // whether this one took effect. + inv.log.Info("operation cancel outcome unknown; not sent again", "error", err) + default: + inv.log.Warn("operation cancel failed", "error", err) + } +} + +// request makes req on the operation until it succeeds or fails for good, +// and returns the last failure. A Busy refusal is repeated after a backoff +// until the invocation halts. When the stream ends, req is repeated on the +// next stream if it had no effect, or if it is idempotent. +func (inv *invocation) request(h handle, idempotent bool, req func(handle) error) error { + backoff := minBackoff + for { + err := req(h) + if err == nil { + return nil + } + f := asFailure(err) + switch { + case h.s.ended() && (idempotent || f.Effect == sandboxwire.EffectNone): + var ok bool + if h, ok = inv.relink(h); !ok { + return err + } + case refused(f): + if !inv.sleep(backoff) { + return err + } + backoff = min(2*backoff, maxBackoff) + default: + return err + } + } +} + +// refused reports whether the service refused a request with Busy, without +// effect; the same request may follow. +func refused(f *sp.Failure) bool { + return f.Code == sp.CodeBusy && f.Effect == sandboxwire.EffectNone +} + +// receive takes one relay message for the invocation. It never blocks, and +// returns an error for a message that breaks the IPC. +func (inv *invocation) receive(m processshim.RelayMessage) error { + switch m := m.(type) { + case processshim.Input: + return inv.answered(m, len(m.Data)) + case processshim.InputEnd: + return inv.answered(m, 0) + case processshim.Written: + w := inv.byFD[m.FD] + if w == nil { + return fmt.Errorf("%w: fd %d has no output", processshim.ErrProtocol, m.FD) + } + if err := w.written(m.Seq); err != nil { + return err + } + inv.acks.deliver(m.Seq) + case processshim.WriteFailed: + w := inv.byFD[m.FD] + if w == nil { + return fmt.Errorf("%w: fd %d has no output", processshim.ErrProtocol, m.FD) + } + sent, err := w.failed(m.Seq) + if err != nil { + return err + } + // The reader is gone; the remote writer gets EPIPE as it would + // locally. The output is delivered at once: closing the remote + // output may wait for a new stream, which the operation's + // settlement may in turn wait behind. + for _, seq := range sent { + inv.acks.deliver(seq) + } + inv.log.Info("output reader gone", "stream", w.stream, "error", unix.Errno(m.Errno)) + inv.helpers.Add(1) + go func() { + defer inv.helpers.Done() + inv.closeOutput(w.stream) + }() + case processshim.Signaled: + select { + case inv.sigs <- m: + default: + inv.log.Warn("signal dropped: too many pending", "signal", m.Number) + } + case processshim.Gone: + inv.helpers.Add(1) + go func() { + defer inv.helpers.Done() + inv.shimGone() + }() + } + return nil +} + +// answered takes the relay's answer to the outstanding Read. +func (inv *invocation) answered(m processshim.RelayMessage, n int) error { + inv.mu.Lock() + credit := inv.credit + inv.credit = 0 + inv.mu.Unlock() + if credit == 0 || n > int(credit) { + return fmt.Errorf("%w: %d bytes of stdin without a Read for them", processshim.ErrProtocol, n) + } + select { + case inv.input <- m: + default: // the stdin pump took the previous answer before granting + } + return nil +} + +// send sends m unless the invocation has ended. +func (inv *invocation) send(m processshim.BrokerMessage) error { + inv.sendMu.Lock() + defer inv.sendMu.Unlock() + if inv.ended { + return errEnded + } + return inv.b.send(m) +} + +// halting reports whether the invocation halted. +func (inv *invocation) halting() bool { + select { + case <-inv.halt: + return true + default: + return false + } +} + +// halted stops every pump and wait. +func (inv *invocation) halted() { + inv.haltOnce.Do(func() { + close(inv.halt) + inv.stopInput() + }) +} + +// stopInput ends stdin forwarding. +func (inv *invocation) stopInput() { + inv.stopOnce.Do(func() { + close(inv.stopIn) + inv.send(processshim.StopInput{ID: inv.rid}) + }) +} + +// teardown ends the invocation: once no relay message reaches it and its +// helpers are done, End tells the relay to close its descriptors. The stdin +// pump is not waited for, so teardown never waits for a stdin request still +// on the stream; it sends nothing after End. +func (inv *invocation) teardown() { + inv.halted() + inv.b.unregister(inv.rid) + inv.writing.Wait() + inv.helpers.Wait() + inv.sendMu.Lock() + defer inv.sendMu.Unlock() + inv.ended = true + inv.b.send(processshim.End{ID: inv.rid}) +} + +func asFailure(err error) *sp.Failure { + var f *sp.Failure + if errors.As(err, &f) { + return f + } + return &sp.Failure{Code: sp.CodeUnknown, Effect: sandboxwire.EffectPossible, Message: err.Error()} +} diff --git a/apps/daemon/internal/processbroker/link_linux.go b/apps/daemon/internal/processbroker/link_linux.go new file mode 100644 index 00000000..b4abd56f --- /dev/null +++ b/apps/daemon/internal/processbroker/link_linux.go @@ -0,0 +1,111 @@ +//go:build linux + +package processbroker + +import ( + "context" + "io" + "log/slog" + "sync" + "time" + + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// stream is one Process stream and the service incarnation behind it. +type stream struct { + client *sp.Client + instance sandboxwire.ID + caps sp.Capabilities +} + +// ended reports whether the stream is over; a request on it fails. +func (s *stream) ended() bool { return s.client.Err() != nil } + +// link keeps one Process stream for the Session and redials it when it ends. +type link struct { + dial func(context.Context) (io.ReadWriteCloser, error) + log *slog.Logger + + // dialing admits one redial at a time; the others wait for its stream. + dialing chan struct{} + + mu sync.Mutex + cur *stream +} + +const ( + minBackoff = 50 * time.Millisecond + maxBackoff = 5 * time.Second +) + +func newLink(dial func(context.Context) (io.ReadWriteCloser, error), log *slog.Logger) *link { + return &link{dial: dial, log: log, dialing: make(chan struct{}, 1)} +} + +// get returns the live stream, redialing until one is up or ctx ends. +func (l *link) get(ctx context.Context) (*stream, error) { + if s := l.current(); s != nil { + return s, nil + } + select { + case l.dialing <- struct{}{}: + case <-ctx.Done(): + return nil, ctx.Err() + } + defer func() { <-l.dialing }() + if s := l.current(); s != nil { + return s, nil + } + backoff := minBackoff + for { + s, err := l.connect(ctx) + if err == nil { + l.mu.Lock() + l.cur = s + l.mu.Unlock() + return s, nil + } + l.log.Warn("process stream unavailable", "error", err, "retry", backoff) + select { + case <-time.After(backoff): + case <-ctx.Done(): + return nil, ctx.Err() + } + backoff = min(2*backoff, maxBackoff) + } +} + +func (l *link) current() *stream { + l.mu.Lock() + defer l.mu.Unlock() + if l.cur != nil && !l.cur.ended() { + return l.cur + } + return nil +} + +// connect dials and describes the service, which pins the incarnation the +// stream's operations belong to. +func (l *link) connect(ctx context.Context) (*stream, error) { + rw, err := l.dial(ctx) + if err != nil { + return nil, err + } + c := sp.NewClient(rw) + d, err := c.Describe(ctx) + if err != nil { + c.Close() + return nil, err + } + return &stream{client: c, instance: d.ServerInstanceID, caps: d.Capabilities}, nil +} + +func (l *link) close() { + l.mu.Lock() + defer l.mu.Unlock() + if l.cur != nil { + l.cur.client.Close() + } +} diff --git a/apps/daemon/internal/processbroker/pumps_linux.go b/apps/daemon/internal/processbroker/pumps_linux.go new file mode 100644 index 00000000..16e85cf9 --- /dev/null +++ b/apps/daemon/internal/processbroker/pumps_linux.go @@ -0,0 +1,442 @@ +//go:build linux + +package processbroker + +import ( + "fmt" + "slices" + "sync" + + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// tracker records which events reached their destination. AckEvents +// follows its contiguous prefix. +type tracker struct { + mu sync.Mutex + prefix uint64 + done map[uint64]bool + changed chan struct{} +} + +func (t *tracker) init() { + t.done = map[uint64]bool{} + t.changed = make(chan struct{}) +} + +func (t *tracker) deliver(seq uint64) { + t.mu.Lock() + defer t.mu.Unlock() + if seq <= t.prefix { + return + } + t.done[seq] = true + start := t.prefix + for t.done[t.prefix+1] { + delete(t.done, t.prefix+1) + t.prefix++ + } + if t.prefix != start { + close(t.changed) + t.changed = make(chan struct{}) + } +} + +func (t *tracker) state() (uint64, <-chan struct{}) { + t.mu.Lock() + defer t.mu.Unlock() + return t.prefix, t.changed +} + +// wait waits until every event through seq is delivered. +func (t *tracker) wait(seq uint64, halt <-chan struct{}) bool { + for { + prefix, changed := t.state() + if prefix >= seq { + return true + } + select { + case <-changed: + case <-halt: + return false + } + } +} + +// chunk is an output event for a writer: data, or the stream's end. +type chunk struct { + seq uint64 + data []byte + close bool +} + +// sent is Output or a Close the relay has not reported yet. +type sent struct { + seq uint64 + n int +} + +// writer forwards one remote stream to its descriptor through the relay. +// At most processshim.OutputWindow bytes are unreported, so a descriptor the +// Harness does not read holds the stream back, and the process service in +// turn holds back the program. The relay closes the descriptor after the +// stream's last byte; on a PTY it also closes descriptor 2, which the merged +// stream never uses. +type writer struct { + inv *invocation + stream sp.Stream + fd uint8 + + mu sync.Mutex + queue []chunk + wake chan struct{} // a push, or window space + window int + sent []sent + broken bool // the relay reported a failed write + last uint64 // the last pushed chunk the relay reports +} + +func (w *writer) push(c chunk) { + w.mu.Lock() + w.queue = append(w.queue, c) + if c.close || len(c.data) > 0 { + w.last = c.seq + } + w.mu.Unlock() + w.signal() +} + +func (w *writer) signal() { + select { + case w.wake <- struct{}{}: + default: + } +} + +// next returns the next chunk once the window has room for it, and whether +// the relay is to report it. +func (w *writer) next() (c chunk, report, ok bool) { + for { + w.mu.Lock() + if len(w.queue) > 0 { + c := w.queue[0] + n := len(c.data) + if w.broken || c.close || w.window == 0 || w.window+n <= processshim.OutputWindow { + w.queue[0] = chunk{} + w.queue = w.queue[1:] + report := !w.broken && (c.close || n > 0) + if report { + w.sent = append(w.sent, sent{seq: c.seq, n: n}) + w.window += n + } + w.mu.Unlock() + return c, report, true + } + } + w.mu.Unlock() + select { + case <-w.wake: + case <-w.inv.halt: + return chunk{}, false, false + } + } +} + +func (w *writer) run() { + inv := w.inv + defer inv.writing.Done() + for { + c, report, ok := w.next() + if !ok { + return + } + switch { + case c.close: + // After a failed write the relay still closes the descriptor, + // without a report. + inv.send(processshim.Close{ID: inv.rid, FD: w.fd, Seq: c.seq}) + if !report { + inv.acks.deliver(c.seq) + } + return + case report: + inv.send(processshim.Output{ID: inv.rid, FD: w.fd, Seq: c.seq, Data: c.data}) + default: + inv.acks.deliver(c.seq) + } + } +} + +// marks returns, for each writer the relay still reports for, the last +// chunk pushed so far. +func (inv *invocation) marks() []processshim.Mark { + var marks []processshim.Mark + for _, w := range inv.byFD { + if w == nil { + continue + } + w.mu.Lock() + if !w.broken && w.last > 0 { + marks = append(marks, processshim.Mark{FD: w.fd, Seq: w.last}) + } + w.mu.Unlock() + } + return marks +} + +// written takes the relay's report of the oldest unreported chunk. +func (w *writer) written(seq uint64) error { + w.mu.Lock() + defer w.mu.Unlock() + if len(w.sent) == 0 || w.sent[0].seq != seq { + return fmt.Errorf("%w: fd %d reported %d out of order", processshim.ErrProtocol, w.fd, seq) + } + w.window -= w.sent[0].n + w.sent = w.sent[1:] + w.signal() + return nil +} + +// failed takes the relay's report that the oldest unreported chunk failed, +// and returns every unreported chunk, none of which the relay reports. +func (w *writer) failed(seq uint64) ([]uint64, error) { + w.mu.Lock() + defer w.mu.Unlock() + if len(w.sent) == 0 || w.sent[0].seq != seq { + return nil, fmt.Errorf("%w: fd %d reported %d out of order", processshim.ErrProtocol, w.fd, seq) + } + seqs := make([]uint64, len(w.sent)) + for i, s := range w.sent { + seqs[i] = s.seq + } + w.sent, w.window, w.broken = nil, 0, true + w.signal() + return seqs, nil +} + +func (inv *invocation) closeOutput(stream sp.Stream) { + inv.request(inv.current(), true, func(h handle) error { return h.op.CloseOutput(inv.b.ctx, stream) }) +} + +// pumpStdin forwards stdin until end of file, the exit, the shim's loss or +// the end. The relay reads descriptor 0 once per Read, so stdin is taken +// only as fast as the service accepts it. +func (inv *invocation) pumpStdin(caps sp.Capabilities) { + limit := min(caps.MaxDataBytes, sandboxwire.MaxChunk) + for { + if !inv.grant(limit) { + inv.stdinStopped() + return + } + var m processshim.RelayMessage + select { + case m = <-inv.input: + case <-inv.stopIn: + inv.stdinStopped() + return + } + select { + case <-inv.stopIn: + inv.stdinStopped() + return + default: + } + in, ok := m.(processshim.Input) + if !ok { // end of file + if inv.term == nil { + inv.closeStdin() + } + return + } + if !inv.writeStdin(in.Data) { + inv.stopInput() + return + } + } +} + +// grant asks the relay for one read of stdin. +func (inv *invocation) grant(limit uint32) bool { + select { + case <-inv.stopIn: + return false + default: + } + inv.mu.Lock() + inv.credit = limit + inv.mu.Unlock() + return inv.send(processshim.Read{ID: inv.rid, Max: limit}) == nil +} + +// stdinStopped closes the remote stdin after the exit, so remote background +// readers see end of file rather than wait for input the Harness no longer +// sends. +func (inv *invocation) stdinStopped() { + if inv.term == nil && inv.exitDecided() { + inv.closeStdin() + } +} + +// writeStdin writes data, which begins at the tracked stdin offset. A Busy +// refusal is repeated after a backoff from the first byte the service did +// not accept. After a lost stream, a failure that may have taken effect or +// an offset conflict, the accepted offset is uncertain: the broker learns it +// with Inspect, on the next stream when the stream ended, and writes only +// the bytes of data after it. WriteStdin succeeds only at the current +// offset, so no byte reaches the program twice. When Inspect shows that the +// leader exited, even before its Exited event arrives, pipe stdin closes at +// that offset, as it does after the exit: a background process may still +// read it, and gets end of file. When the offset cannot be +// learned, lies outside data, or shows that a stream that is still up +// failed the write without taking it, the program would wait for input +// that never comes: stdinLost ends the invocation. +func (inv *invocation) writeStdin(data []byte) bool { + h := inv.current() + at := h.op.StdinOffset() // data[0]'s offset + backoff := minBackoff + for { + n, err := h.op.WriteStdin(inv.b.ctx, data) + if err == nil { + return true + } + f, lost := asFailure(err), h.s.ended() + switch { + case inv.b.ctx.Err() != nil: + return false + case !lost && refused(f): + data, at = data[n:], at+uint64(n) // the offset already counts the accepted n + if !inv.sleep(backoff) { + return false + } + backoff = min(2*backoff, maxBackoff) + continue + case !lost && f.Effect == sandboxwire.EffectNone && f.Code != sp.CodeInputOffsetConflict: + if f.Code != sp.CodeStdinClosed && f.Code != sp.CodeNotRunning && f.Code != sp.CodeReleased { + inv.log.Warn("stdin forwarding stopped", "error", err) + } + return false + } + var st sp.OperationStatus + err = inv.request(h, true, func(next handle) error { + h = next + var err error + st, err = h.op.Inspect(inv.b.ctx) + return err + }) + off := st.StdinOffset + switch { + case err != nil && (inv.b.ctx.Err() != nil || inv.halting()): + return false + case err != nil: + inv.stdinLost(fmt.Sprintf("stdin could not be resumed: %s", asFailure(err).Message)) + return false + case st.StdinClosed: + return false + case st.State != sp.StateRunning: + if inv.term == nil { + inv.request(h, true, func(h handle) error { return h.op.CloseStdin(inv.b.ctx) }) + } + return false + case off < at || off-at > uint64(len(data)): + inv.stdinLost(fmt.Sprintf("stdin could not be resumed: the service accepted %d bytes, outside %d to %d", off, at, at+uint64(len(data)))) + return false + case !lost && off == at+uint64(n): + inv.stdinLost(fmt.Sprintf("stdin could not be resumed: %s", f.Message)) + return false + } + data, at = data[off-at:], off + if len(data) == 0 { + return true + } + } +} + +// stdinLost ends an invocation whose stdin cannot continue: unless the exit +// is decided, the shim exits with 255 and the reason, and the program, which +// would wait for input that never comes, is cancelled. When the shim already +// has its Result or is gone, the path that answered or lost it owns the +// program's end, and stdinLost cancels nothing. +func (inv *invocation) stdinLost(reason string) { + if !inv.replyBeforeExit(processshim.Result{Code: processshim.ExitLost, Message: message(reason)}) { + return + } + inv.log.Warn("process invocation failed", "reason", reason) + inv.cancelRemote() +} + +func (inv *invocation) closeStdin() { + inv.request(inv.current(), true, func(h handle) error { return h.op.CloseStdin(inv.b.ctx) }) +} + +func (inv *invocation) exitDecided() bool { + inv.mu.Lock() + defer inv.mu.Unlock() + return inv.exited +} + +// ackLoop acknowledges each delivered prefix, letting the service reclaim +// its replay and keep reading output. +func (inv *invocation) ackLoop() { + defer inv.helpers.Done() + var acked uint64 + for { + prefix, changed := inv.acks.state() + if prefix > acked { + if err := inv.request(inv.current(), true, func(h handle) error { return h.op.Ack(inv.b.ctx, prefix) }); err != nil { + return // released, or the operation is gone + } + acked = prefix + continue + } + select { + case <-changed: + case <-inv.halt: + return + } + } +} + +func (inv *invocation) forwardSignals() { + defer inv.helpers.Done() + for { + select { + case n := <-inv.sigs: + inv.signal(n) + case <-inv.halt: + return + } + } +} + +// signal forwards one signal the shim caught. A signal whose delivery is +// uncertain is not resent. +func (inv *invocation) signal(s processshim.Signaled) { + h, n := inv.current(), s.Number + if inv.term != nil && n == uint16(unix.SIGWINCH) { + if s.Size != nil { + size := windowSize(*s.Size) + inv.request(h, true, func(h handle) error { return h.op.Resize(inv.b.ctx, size) }) + } + return + } + target := sp.TargetInitialProcessGroup + if inv.term != nil && slices.Contains(ptyGroupSignals, n) { + target = sp.TargetPTYForegroundGroup + } + sig := sp.Signal(n) + if f := h.s.caps.CheckSignal(sig, target); f != nil { + inv.log.Info("signal not forwarded", "signal", n, "reason", f.Message) + return + } + err := inv.request(h, false, func(h handle) error { return h.op.Signal(inv.b.ctx, sig, target) }) + if err == nil { + return + } + if f := asFailure(err); f.Code != sp.CodeNotRunning && f.Code != sp.CodeReleased { + inv.log.Info("signal not delivered", "signal", n, "error", err) + } +} diff --git a/apps/daemon/internal/processbroker/view_linux_test.go b/apps/daemon/internal/processbroker/view_linux_test.go new file mode 100644 index 00000000..8cb34997 --- /dev/null +++ b/apps/daemon/internal/processbroker/view_linux_test.go @@ -0,0 +1,224 @@ +//go:build linux + +package processbroker + +import ( + "context" + "fmt" + "io" + "net" + "os" + "path/filepath" + "sync" + "syscall" + "testing" + "time" + + gofs "github.com/hanwen/go-fuse/v2/fs" + "github.com/hanwen/go-fuse/v2/fuse" + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/sessionview" + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" +) + +// harnessEnv makes the test binary the Harness of TestStuckOutputDoesNotHoldTeardown. +const harnessEnv = "OAC_TEST_VIEW_HARNESS" + +func TestViewRunsRemoteShell(t *testing.T) { + w := newHangWorld(t) + v, _ := startView(t, w, sessionview.Process{Path: "/bin/sh", Args: []string{"sh", "-c", "echo $0"}}) + out, err := io.ReadAll(v.Stdout()) + if err != nil { + t.Fatal(err) + } + if exit, err := v.Wait(); err != nil || exit != (sessionview.Exit{}) || string(out) != "sh\n" { + t.Fatalf("Wait = %+v, %v; output %q", exit, err, out) + } +} + +// TestStuckOutputDoesNotHoldTeardown gives a shim's stdout on a world file +// whose writes are never answered. Broker.Close returns while the relay's +// write is stuck, and the view's teardown ends it by stopping the world, as +// the world frontend's Stop does. +func TestStuckOutputDoesNotHoldTeardown(t *testing.T) { + w := newHangWorld(t) + v, b := startView(t, w, sessionview.Process{Path: "/.oac/harness/harness", Args: []string{"harness"}, Env: []string{harnessEnv + "=1"}}) + select { + case <-w.hung: + case <-time.After(20 * time.Second): + t.Fatal("the relay never wrote the output") + } + started := time.Now() + b.Close() + if elapsed := time.Since(started); elapsed > 5*time.Second { + t.Errorf("Broker.Close took %v while the relay's write was stuck", elapsed) + } + started = time.Now() + if err := v.Close(); err != nil { + t.Errorf("View.Close = %v", err) + } + if elapsed := time.Since(started); elapsed > 20*time.Second { + t.Errorf("View.Close took %v", elapsed) + } +} + +// runHarness runs the shim with its stdout on the world file whose writes +// are never answered. +func runHarness() int { + f, err := os.OpenFile("/hang", os.O_WRONLY, 0) + if err == nil { + err = unix.Dup3(int(f.Fd()), 1, 0) + } + if err == nil { + err = syscall.Exec("/bin/sh", []string{"sh", "-c", "echo hi"}, []string{}) + } + fmt.Fprintln(os.Stderr, err) + return 1 +} + +// startView starts a view whose process runs as viewID with the shim at +// /bin/sh, and a broker for its relay. The test binary is at +// /.oac/harness/harness. +func startView(t *testing.T, w *hangWorld, p sessionview.Process) (*sessionview.View, *Broker) { + if os.Getenv(viewGateEnv) != "1" { + t.Skipf("set %s=1 and run the test binary as root in a privileged container; see the comment at the top of broker_linux_test.go", viewGateEnv) + } + if err := sessionview.Probe(); err != nil { + t.Fatalf("Probe: %v", err) + } + sock := service.socket(t) + harness := t.TempDir() + self, err := os.Executable() + if err != nil { + t.Fatal(err) + } + data, err := os.ReadFile(self) + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(harness, "harness"), data, 0o755); err != nil { + t.Fatal(err) + } + p.Dir, p.UID, p.GID, p.Stderr = "/", viewID, viewID, os.Stderr + v, err := sessionview.Start(context.Background(), sessionview.Spec{ + World: w.serve, + Private: []sessionview.PrivateDir{{Name: "harness", HostDir: harness, Exec: true}}, + Shim: sessionview.Shim{Binary: shimBinary(t), Paths: []string{"/bin/sh"}}, + Process: p, + StagingParent: t.TempDir(), + }) + if err != nil { + t.Fatalf("Start: %v", err) + } + t.Cleanup(func() { v.Close() }) + b, err := Start(Config{ + Relay: v.Relay(), + Executables: Executables{Paths: map[string]string{"/bin/sh": "/bin/sh"}}, + Environment: Environment{Sandbox: map[string]string{"PATH": "/usr/bin:/bin"}}, + Scope: sp.ScopePOSIXSession, + Dial: func(ctx context.Context) (io.ReadWriteCloser, error) { + return new(net.Dialer).DialContext(ctx, "unix", sock) + }, + CancelGrace: time.Second, + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { b.Close() }) + return v, b +} + +// shimBinary returns a static oac-process-shim for the view. +func shimBinary(t *testing.T) string { + if bin := os.Getenv(shimEnv); bin != "" { + return bin + } + bin := filepath.Join(t.TempDir(), "oac-process-shim") + if err := build(bin, "../../cmd/oac-process-shim"); err != nil { + t.Fatal(err) + } + return bin +} + +// hangWorld serves a directory as the world, the way the world frontend +// serves a sandbox, presenting each mountpoint at its view path. Writes to +// /hang get no answer, even when the writer is killed, until Stop: like the +// world frontend's, Stop ends every request still pending and then waits for +// serving to end. +type hangWorld struct { + dir string + served chan struct{} + hung, release chan struct{} // hung closes at the first write to /hang + hungOnce sync.Once + releaseOnce sync.Once +} + +func newHangWorld(t *testing.T) *hangWorld { + dir := t.TempDir() + for _, d := range []string{".oac/harness", ".oac/run", ".oac/bin", "proc", "dev", "bin"} { + if err := os.MkdirAll(filepath.Join(dir, d), 0o755); err != nil { + t.Fatal(err) + } + } + for _, f := range []string{"bin/sh", "hang"} { + if err := os.WriteFile(filepath.Join(dir, f), nil, 0o666); err != nil { + t.Fatal(err) + } + } + return &hangWorld{dir: dir, hung: make(chan struct{}), release: make(chan struct{})} +} + +func (w *hangWorld) serve(_ context.Context, dev *os.File, mount sessionview.WorldMount) (sessionview.WorldServer, sessionview.Presentation, error) { + var p sessionview.Presentation + root, err := gofs.NewLoopbackRoot(w.dir) + if err != nil { + return nil, p, err + } + root.(*gofs.LoopbackNode).RootData.NewNode = func(r *gofs.LoopbackRoot, parent *gofs.Inode, name string, _ *syscall.Stat_t) gofs.InodeEmbedder { + n := &gofs.LoopbackNode{RootData: r} + if name == "hang" && parent.IsRoot() { + return &hangNode{LoopbackNode: n, w: w} + } + return n + } + fd, err := unix.Dup(int(dev.Fd())) + if err != nil { + return nil, p, err + } + srv, err := fuse.NewServer(gofs.NewNodeFS(root, &gofs.Options{}), fmt.Sprintf("/dev/fd/%d", fd), &fuse.MountOptions{}) + if err != nil { + unix.Close(fd) + return nil, p, err + } + w.served = make(chan struct{}) + go func() { + srv.Serve() + close(w.served) + }() + for _, m := range mount.Mountpoints { + p.Targets = append(p.Targets, m.Path) + } + return w, p, nil +} + +func (w *hangWorld) Stop() error { + w.releaseOnce.Do(func() { close(w.release) }) + select { + case <-w.served: + return nil + case <-time.After(10 * time.Second): + return fmt.Errorf("world still serving 10s after its requests ended") + } +} + +type hangNode struct { + *gofs.LoopbackNode + w *hangWorld +} + +func (n *hangNode) Write(context.Context, gofs.FileHandle, []byte, int64) (uint32, syscall.Errno) { + n.w.hungOnce.Do(func() { close(n.w.hung) }) + <-n.w.release + return 0, syscall.EIO +} diff --git a/apps/daemon/internal/processshim/conn_linux.go b/apps/daemon/internal/processshim/conn_linux.go new file mode 100644 index 00000000..bf92b442 --- /dev/null +++ b/apps/daemon/internal/processshim/conn_linux.go @@ -0,0 +1,156 @@ +//go:build linux + +package processshim + +import ( + "errors" + "fmt" + "io" + "net" + + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// Conn is one end of an invocation connection. +type Conn struct { + c *net.UnixConn + // fds collects descriptors while the relay reads the Request; nil + // otherwise, so any control data is a violation. + fds *[]int + // first is true until the first read returns. + first bool +} + +// NewConn wraps a connected Unix stream socket. +func NewConn(c *net.UnixConn) *Conn { return &Conn{c: c, first: true} } + +// Close closes the socket. +func (c *Conn) Close() error { return c.c.Close() } + +// Unix returns the socket. +func (c *Conn) Unix() *net.UnixConn { return c.c } + +// SendRequest sends r with fds attached to its first byte. +func (c *Conn) SendRequest(r Request, fds [3]int) error { + f := Frame(r) + if len(f.Payload) > MaxRequestBytes { + return fmt.Errorf("%w: request of %d bytes exceeds %d", ErrProtocol, len(f.Payload), MaxRequestBytes) + } + var b frameBuffer + if err := sandboxwire.WriteFrame(&b, f); err != nil { + return err + } + n, _, err := c.c.WriteMsgUnix(b, unix.UnixRights(fds[:]...), nil) + if err != nil { + return err + } + _, err = c.c.Write(b[n:]) + return err +} + +// Send sends a message without descriptors. +func (c *Conn) Send(m Message) error { return sandboxwire.WriteFrame(c.c, Frame(m)) } + +// ReadRequest reads the Request and the three descriptors attached to it. The +// caller owns the descriptors once it returns without error; on error none +// remain open. A Request with another Version returns with only its Version. +func (c *Conn) ReadRequest() (Request, [3]int, error) { + var fds []int + c.fds = &fds + m, err := c.read(MaxRequestBytes) + c.fds = nil + if err == nil && len(fds) != 3 { + err = fmt.Errorf("%w: request carries %d descriptors", ErrProtocol, len(fds)) + } + r, ok := m.(Request) + if err == nil && !ok { + err = fmt.Errorf("%w: first message is type %d", ErrProtocol, Frame(m).Type) + } + if err != nil { + closeAll(fds) + return Request{}, [3]int{-1, -1, -1}, err + } + return r, [3]int(fds), nil +} + +// ReadMessage reads one message. Descriptors on any message other than the +// Request are a violation. +func (c *Conn) ReadMessage() (Message, error) { return c.read(MaxFrameBytes) } + +func (c *Conn) read(max uint32) (Message, error) { + f, err := sandboxwire.ReadFrame(reader{c}, max) + if err != nil { + return nil, err + } + return Decode(f) +} + +// reader reads the socket for ReadFrame and checks the control data that +// arrives with each read. +type reader struct{ *Conn } + +func (c reader) Read(p []byte) (int, error) { + oob := make([]byte, unix.CmsgSpace(3*4)) + n, oobn, flags, _, err := c.c.ReadMsgUnix(p, oob) + first := c.first + c.first = false + if oobn == 0 && flags&unix.MSG_CTRUNC == 0 { + if err == nil && n == 0 && len(p) > 0 { + err = io.EOF + } + return n, err + } + fds, perr := parseRights(oob[:oobn]) + switch { + case flags&unix.MSG_CTRUNC != 0: + perr = errors.New("truncated control data") + case perr == nil && (c.fds == nil || !first): + perr = errors.New("unexpected descriptors") + } + if perr != nil { + closeAll(fds) + return 0, fmt.Errorf("%w: %w", ErrProtocol, perr) + } + *c.fds = fds + return n, err +} + +// parseRights returns the descriptors in one SCM_RIGHTS message. It returns +// every descriptor it found, even with an error, so the caller can close them. +func parseRights(oob []byte) ([]int, error) { + msgs, err := unix.ParseSocketControlMessage(oob) + if err != nil { + return nil, err + } + var fds []int + for _, m := range msgs { + if m.Header.Level != unix.SOL_SOCKET || m.Header.Type != unix.SCM_RIGHTS { + err = fmt.Errorf("control message %d/%d", m.Header.Level, m.Header.Type) + continue + } + got, perr := unix.ParseUnixRights(&m) + fds = append(fds, got...) + if perr != nil { + err = perr + } + } + if err == nil && len(msgs) != 1 { + err = fmt.Errorf("%d control messages", len(msgs)) + } + return fds, err +} + +func closeAll(fds []int) { + for _, fd := range fds { + unix.Close(fd) + } +} + +type frameBuffer []byte + +func (b *frameBuffer) Write(p []byte) (int, error) { + *b = append(*b, p...) + return len(p), nil +} diff --git a/apps/daemon/internal/processshim/conn_linux_test.go b/apps/daemon/internal/processshim/conn_linux_test.go new file mode 100644 index 00000000..b2e654de --- /dev/null +++ b/apps/daemon/internal/processshim/conn_linux_test.go @@ -0,0 +1,107 @@ +package processshim + +import ( + "errors" + "net" + "os" + "path/filepath" + "reflect" + "testing" + + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +func connPair(t *testing.T) (*Conn, *Conn) { + t.Helper() + ln, err := net.ListenUnix("unix", &net.UnixAddr{Name: filepath.Join(t.TempDir(), "s"), Net: "unix"}) + if err != nil { + t.Fatal(err) + } + defer ln.Close() + client, err := net.DialUnix("unix", nil, ln.Addr().(*net.UnixAddr)) + if err != nil { + t.Fatal(err) + } + server, err := ln.AcceptUnix() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { client.Close(); server.Close() }) + return NewConn(client), NewConn(server) +} + +func pipeFDs(t *testing.T) ([3]int, *os.File) { + t.Helper() + r, w, err := os.Pipe() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { r.Close(); w.Close() }) + fd := int(w.Fd()) + return [3]int{fd, fd, fd}, r +} + +func TestRequestCarriesCloseOnExecDescriptors(t *testing.T) { + shim, broker := connPair(t) + fds, r := pipeFDs(t) + req := fixtures[0].msg.(Request) + if err := shim.SendRequest(req, fds); err != nil { + t.Fatal(err) + } + got, recv, err := broker.ReadRequest() + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(got, req) { + t.Fatalf("request %#v, want %#v", got, req) + } + for _, fd := range recv { + flags, err := unix.FcntlInt(uintptr(fd), unix.F_GETFD, 0) + if err != nil || flags&unix.FD_CLOEXEC == 0 { + t.Fatalf("fd %d flags %#x: %v", fd, flags, err) + } + } + if _, err := unix.Write(recv[1], []byte("x")); err != nil { + t.Fatal(err) + } + buf := make([]byte, 1) + if _, err := r.Read(buf); err != nil || buf[0] != 'x' { + t.Fatalf("read %q: %v", buf, err) + } + closeAll(recv[:]) +} + +func TestUnexpectedDescriptorsAreRejected(t *testing.T) { + t.Run("request without descriptors", func(t *testing.T) { + shim, broker := connPair(t) + if err := shim.Send(fixtures[0].msg); err != nil { + t.Fatal(err) + } + if _, _, err := broker.ReadRequest(); !errors.Is(err, ErrProtocol) { + t.Fatalf("got %v, want ErrProtocol", err) + } + }) + t.Run("descriptors on a later message", func(t *testing.T) { + shim, broker := connPair(t) + fds, _ := pipeFDs(t) + if err := shim.SendRequest(fixtures[0].msg.(Request), fds); err != nil { + t.Fatal(err) + } + _, recv, err := broker.ReadRequest() + if err != nil { + t.Fatal(err) + } + closeAll(recv[:]) + frame := Frame(Signal{Number: 2}) + var buf frameBuffer + _ = sandboxwire.WriteFrame(&buf, frame) + if _, _, err := shim.c.WriteMsgUnix(buf, unix.UnixRights(fds[0]), nil); err != nil { + t.Fatal(err) + } + if _, err := broker.ReadMessage(); !errors.Is(err, ErrProtocol) { + t.Fatalf("got %v, want ErrProtocol", err) + } + }) +} diff --git a/apps/daemon/internal/processshim/endpoint_linux.go b/apps/daemon/internal/processshim/endpoint_linux.go new file mode 100644 index 00000000..284ec369 --- /dev/null +++ b/apps/daemon/internal/processshim/endpoint_linux.go @@ -0,0 +1,272 @@ +//go:build linux + +package processshim + +import ( + "encoding/binary" + "errors" + "fmt" + "sync" + "unsafe" + + "golang.org/x/sys/unix" +) + +// pipeBuf is PIPE_BUF: a write of at most this many bytes to a pipe that +// poll reports writable does not block. +const pipeBuf = 4096 + +const ( + ptySlaveMajor = 136 // UNIX98_PTY_SLAVE_MAJOR + ptySlaveMajors = 8 // UNIX98_PTY_MAJOR_COUNT +) + +var errStopped = errors.New("stopped") + +// stopFlag is a level-triggered flag that pollers include, with a channel +// that closes when it is set. Once set it stays set. Setting it after close +// does nothing, so a late set never writes to a reused descriptor number. +type stopFlag struct { + fd int + c chan struct{} + + mu sync.Mutex + raised, done bool +} + +func newStopFlag() (*stopFlag, error) { + fd, err := unix.Eventfd(0, unix.EFD_CLOEXEC|unix.EFD_NONBLOCK) + if err != nil { + return nil, err + } + return &stopFlag{fd: fd, c: make(chan struct{})}, nil +} + +func (s *stopFlag) set() { + s.mu.Lock() + defer s.mu.Unlock() + if !s.raised { + s.raised = true + close(s.c) + if !s.done { + var one [8]byte + binary.NativeEndian.PutUint64(one[:], 1) + unix.Write(s.fd, one[:]) + } + } +} + +func (s *stopFlag) isSet() bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.raised +} + +// close closes the flag's descriptor once its pollers have returned. +func (s *stopFlag) close() { + s.mu.Lock() + defer s.mu.Unlock() + if !s.done { + s.done = true + unix.Close(s.fd) + } +} + +type ioKind uint8 + +const ( + ioOwn ioKind = iota // the relay's own non-blocking description + ioSocket // a socket, used with MSG_DONTWAIT + ioFile // a regular file or block device, used as it is + ioShared // another description, polled before each call +) + +// endpoint is a received descriptor and the descriptor its I/O uses. +type endpoint struct { + fd int // the received descriptor + io int // fd, an own description, or -1 for a FIFO without a reader + kind ioKind +} + +// openEndpoint prepares received descriptor fd for reading or writing. A +// pipe, FIFO or pty slave is reopened through /proc/self/fd as the relay's +// own non-blocking description, so a poll that ending the invocation +// interrupts is the only wait on a peer; a socket gets MSG_DONTWAIT. A +// regular file or block device is used as it is. Any other descriptor, and +// one the relay may not reopen, is shared with the Harness: the relay polls +// it before each call and writes at most pipeBuf bytes at once. The flags of +// the shared description never change. +func openEndpoint(fd int, write bool) (endpoint, error) { + var st unix.Stat_t + if err := unix.Fstat(fd, &st); err != nil { + return endpoint{}, err + } + shared := endpoint{fd: fd, io: fd, kind: ioShared} + switch typ := st.Mode & unix.S_IFMT; typ { + case unix.S_IFREG, unix.S_IFBLK: + return endpoint{fd: fd, io: fd, kind: ioFile}, nil + case unix.S_IFSOCK: + return endpoint{fd: fd, io: fd, kind: ioSocket}, nil + case unix.S_IFCHR, unix.S_IFIFO: + if major := unix.Major(st.Rdev); typ == unix.S_IFCHR && (major < ptySlaveMajor || major >= ptySlaveMajor+ptySlaveMajors) { + return shared, nil + } + mode, use := unix.O_RDONLY, "reading" + if write { + mode, use = unix.O_WRONLY, "writing" + } + // The own description gets no access the received one lacks. + fl, err := unix.FcntlInt(uintptr(fd), unix.F_GETFL, 0) + if err != nil { + return endpoint{}, err + } + if acc := fl & unix.O_ACCMODE; fl&unix.O_PATH != 0 || acc != mode && acc != unix.O_RDWR { + return endpoint{}, errors.New("not open for " + use) + } + io, err := unix.Open(fmt.Sprintf("/proc/self/fd/%d", fd), mode|unix.O_NONBLOCK|unix.O_NOCTTY|unix.O_CLOEXEC, 0) + switch { + case err == unix.ENXIO && write && typ == unix.S_IFIFO: + return endpoint{fd: fd, io: -1, kind: ioOwn}, nil // no reader + case err == unix.EACCES || err == unix.EPERM: + return shared, nil + case err != nil: + return endpoint{}, fmt.Errorf("reopen: %w", err) + } + return endpoint{fd: fd, io: io, kind: ioOwn}, nil + default: + return endpoint{}, fmt.Errorf("file type %#o is not supported", typ) + } +} + +func (e endpoint) close() { + if e.io >= 0 && e.io != e.fd { + unix.Close(e.io) + } + unix.Close(e.fd) +} + +func (e endpoint) read(buf []byte) (int, error) { + if e.kind == ioSocket { + n, _, errno := unix.Syscall6(unix.SYS_RECVFROM, uintptr(e.io), uintptr(unsafe.Pointer(&buf[0])), uintptr(len(buf)), unix.MSG_DONTWAIT, 0, 0) + if errno != 0 { + return 0, errno + } + return int(n), nil + } + return unix.Read(e.io, buf) +} + +// write writes once. Output on a Unix socket carries the relay's own +// credentials. +func (e endpoint) write(data []byte) (int, error) { + switch e.kind { + case ioSocket: + return unix.SendmsgN(e.io, data, nil, nil, unix.MSG_DONTWAIT|unix.MSG_NOSIGNAL) + case ioShared: + data = data[:min(len(data), pipeBuf)] + } + if e.io < 0 { + return 0, unix.EPIPE + } + return unix.Write(e.io, data) +} + +// readFD reads once. It returns 0 and nil at end of file. +func readFD(e endpoint, buf []byte, stops ...*stopFlag) (int, error) { + for { + if stopped(stops) { + return 0, errStopped + } + if e.kind == ioShared { + if err := waitFD(e.io, unix.POLLIN, stops...); err != nil { + return 0, err + } + } + n, err := e.read(buf) + switch err { + case unix.EINTR: + continue + case unix.EAGAIN: + if err := waitFD(e.io, unix.POLLIN, stops...); err != nil { + return 0, err + } + continue + } + return max(n, 0), err + } +} + +// writeFD writes all of data, handling short writes. +func writeFD(e endpoint, data []byte, stops ...*stopFlag) error { + for len(data) > 0 { + if stopped(stops) { + return errStopped + } + if e.kind == ioShared { + if err := waitFD(e.io, unix.POLLOUT, stops...); err != nil { + return err + } + } + n, err := e.write(data) + switch { + case err == unix.EINTR: + continue + case err == unix.EAGAIN: + if err := waitFD(e.io, unix.POLLOUT, stops...); err != nil { + return err + } + continue + case err != nil: + return err + } + data = data[n:] + } + return nil +} + +// tryWrite writes data only as far as the endpoint takes it now. +func tryWrite(e endpoint, data []byte) { + if e.kind == ioShared { + p := []unix.PollFd{{Fd: int32(e.io), Events: unix.POLLOUT}} + if n, err := unix.Poll(p, 0); n != 1 || err != nil || p[0].Revents&unix.POLLOUT == 0 { + return + } + } + e.write(data[:min(len(data), pipeBuf)]) +} + +func stopped(stops []*stopFlag) bool { + for _, s := range stops { + if s.isSet() { + return true + } + } + return false +} + +// waitFD waits until fd has events or a stop flag is set. +func waitFD(fd int, events int16, stops ...*stopFlag) error { + pfds := []unix.PollFd{{Fd: int32(fd), Events: events}} + for _, s := range stops { + pfds = append(pfds, unix.PollFd{Fd: int32(s.fd), Events: unix.POLLIN}) + } + for { + if _, err := unix.Poll(pfds, -1); err != nil { + if err == unix.EINTR { + continue + } + return err + } + for _, p := range pfds[1:] { + if p.Revents != 0 { + return errStopped + } + } + switch r := pfds[0].Revents; { + case r&unix.POLLNVAL != 0: + return unix.EBADF + case r != 0: + return nil + } + } +} diff --git a/apps/daemon/internal/processshim/endpoint_linux_test.go b/apps/daemon/internal/processshim/endpoint_linux_test.go new file mode 100644 index 00000000..62c2e11d --- /dev/null +++ b/apps/daemon/internal/processshim/endpoint_linux_test.go @@ -0,0 +1,66 @@ +//go:build linux + +package processshim + +import ( + "testing" + "time" + + "golang.org/x/sys/unix" +) + +// Two invocations writing to one pipe race for the space a reader frees; the +// loser must still stop when aborted, and the shared description must keep +// its flags. +func TestWritersSharingAFullPipeStop(t *testing.T) { + var p [2]int + if err := unix.Pipe2(p[:], unix.O_CLOEXEC); err != nil { + t.Fatal(err) + } + defer unix.Close(p[0]) + defer unix.Close(p[1]) + done := make(chan error, 2) + var stops []*stopFlag + for range 2 { + fd, err := unix.FcntlInt(uintptr(p[1]), unix.F_DUPFD_CLOEXEC, 0) + if err != nil { + t.Fatal(err) + } + ep, err := openEndpoint(fd, true) + if err != nil { + t.Fatal(err) + } + defer ep.close() + stop, err := newStopFlag() + if err != nil { + t.Fatal(err) + } + defer stop.close() + stops = append(stops, stop) + go func() { done <- writeFD(ep, make([]byte, 1<<20), stop) }() + } + // Free space now and then so both writers wake for it, then stop reading. + buf := make([]byte, pipeBuf) + for range 20 { + time.Sleep(10 * time.Millisecond) + if _, err := unix.Read(p[0], buf); err != nil { + t.Fatal(err) + } + } + for _, s := range stops { + s.set() + } + for range 2 { + select { + case err := <-done: + if err != errStopped { + t.Fatalf("writeFD = %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("a writer did not stop") + } + } + if fl, err := unix.FcntlInt(uintptr(p[1]), unix.F_GETFL, 0); err != nil || fl&unix.O_NONBLOCK != 0 { + t.Fatalf("shared flags %#x, %v", fl, err) + } +} diff --git a/apps/daemon/internal/processshim/ipc.go b/apps/daemon/internal/processshim/ipc.go new file mode 100644 index 00000000..35799d4e --- /dev/null +++ b/apps/daemon/internal/processshim/ipc.go @@ -0,0 +1,845 @@ +// Package processshim is oac-process-shim, the Session's process relay, and +// their local IPC with the process broker (apps/daemon/internal/processbroker). +// The IPC is private to the shim, the relay and the broker; it is not part of +// the process protocol. +// +// A Harness in a Session view executes the shim under a declared name or path. +// The shim holds no credentials and never contacts the sandbox: it hands its +// invocation to the relay, and the broker runs the program in the sandbox +// through the process protocol. The relay is the same binary in relay mode, +// one per Session, running in the view as the Session user with no +// capabilities. It is the only process that receives or operates on the +// shim's descriptors; the broker, outside the view, holds only its own end of +// the relay connection. +// +// # Shim and relay +// +// One invocation is one connection to SocketPath: +// +// 1. The shim sends a Request. The first byte of its frame carries, as one +// SCM_RIGHTS message, exactly three descriptors: the shim's fds 0, 1 and +// 2, in that order. +// 2. The relay answers with an Ack once the broker accepts the invocation, or +// with a Result instead of the Ack when the invocation is refused. +// 3. After the Ack the shim closes its fds 0, 1 and 2, so the relay's copies +// are the invocation's only references to them, and sends a Signal for +// each signal it catches. +// 4. The relay sends exactly one Result, and the shim exits with it. +// +// The relay owns the received descriptors from the moment it reads them and +// closes each one when it is done with it; after a refusing Result it has +// closed all three. Whoever holds fd 2 writes a failure message: the shim +// before the Ack, the relay after it. Received descriptors are close-on-exec. +// Truncated control data, anything other than one SCM_RIGHTS message with +// three descriptors on the Request, and any control data on a later frame are +// protocol violations; the receiver closes every descriptor it got and drops +// the connection. +// +// # Relay and broker +// +// The relay and the broker share one stream socket, which sessionview creates; +// the relay holds its end at RelayBrokerFD. Each message names its invocation +// by an ID the relay assigns, starting at 1 and increasing. No descriptor +// crosses this connection: the broker reads it without a control buffer, so +// the kernel discards any SCM_RIGHTS the relay attaches. The broker treats +// every relay message as untrusted Session input, and a message that breaks +// these rules ends the connection. +// +// 1. The relay sends Open with the shim's Request and, when fds 0 and 1 are +// both terminals, the Terminal they share. +// 2. The broker answers with Accept, after which the relay sends the shim its +// Ack, or refuses with an Exit whose Result carries the reason. +// 3. The broker sends Started once the program runs; on a terminal the relay +// then makes the terminal raw, and only after that writes its output or +// reads its input. Terminal output that arrives after the Exit is written +// in the restored mode, so the local terminal processes it again (a remote +// "\r\n" becomes "\r\r\n"). +// 4. Stdin is read on demand. Each Read grants one Input or InputEnd; the +// relay reads fd 0 once per grant, sends what it read, and sends InputEnd +// at end of file or on a read error. StopInput ends reading, and the relay +// closes fd 0 without sending InputEnd. +// 5. The broker sends each output stream as Output and then Close, with the +// process protocol's event sequence numbers, on FD 1 (stdout or the +// terminal) or FD 2 (stderr). The relay writes each FD's messages in order +// and reports each one with Written once it is written or, for Close, once +// the descriptor is closed. On a terminal, Close of FD 1 closes fd 2 too. +// The first write that fails is reported with WriteFailed; the relay then +// reports nothing more for that FD, discards its Output and still closes +// it on Close. At most OutputWindow bytes of unreported Output are +// outstanding per FD. +// 6. The relay reports each signal the shim forwards with Signaled, and the +// shim's loss before its Result with Gone. +// 7. Exit ends the shim. The relay stops reading stdin, waits until every +// Mark's FD has reported Written through the Mark's Seq or has failed, +// restores the terminal and sends the shim the Result. A Result Message +// goes to the shim in the Result before the Ack, and to fd 2 after it. +// 8. Notice writes "oac-process-shim: " to fd 2 after the output +// queued before it. +// 9. End closes the invocation: the relay stops its pumps, closes every +// descriptor and forgets the ID. When the shim has had no Result, it gets +// ExitLost. The broker sends nothing with the ID after End; it ignores +// relay messages for an ID it has ended. +// +// The relay has at most MaxInvocations open and refuses the shim's request +// beyond that. When the broker's end closes, the relay writes the reason to +// each invocation's fd 2, sends each waiting shim ExitLost and exits. +// +// Frames use the internal/sandboxwire header with RequestID zero, and payloads +// use its primitive encoding. +package processshim + +import ( + "errors" + "fmt" + "slices" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +const ( + // SocketName is the relay's socket in the view's run directory. + SocketName = "process.sock" + // SocketPath is the fixed path the shim connects to: SocketName in the + // run directory of the view layout in package agent. It is not + // configurable, so the Harness environment cannot redirect it. It is a + // literal so that the shim does not link package agent; a test checks it + // against the layout. + SocketPath = "/.oac/run/" + SocketName + + // RelayName is the relay's name in the view's shim directory, which no + // declared shim takes, and RelayPath is where sessionview executes the + // relay, with argv RelayArgs. Like SocketPath they are literals that a + // test checks against the layout. + RelayName = "oac-process-shim" + RelayPath = "/.oac/bin/" + RelayName + // RelayBrokerFD is the relay's end of the broker connection, and + // RelayListenerFD the listening socket bound at SocketPath, when the + // relay starts. + RelayBrokerFD = 3 + RelayListenerFD = 4 + + // Version is the shim's IPC version. The relay refuses any other. + Version = 1 + // MaxFrameBytes bounds a frame payload. + MaxFrameBytes = sandboxwire.MaxPayload + // MaxRequestBytes bounds a Request's payload, and so the argument list and + // environment, leaving room for Open around it. + MaxRequestBytes = MaxFrameBytes - 1024 + // MaxMessageBytes bounds a Result or Notice message. + MaxMessageBytes = 4096 + // MaxControlChars bounds Terminal.Cc. + MaxControlChars = 32 + // MaxInvocations bounds the invocations a relay has open. + MaxInvocations = 256 + // OutputWindow bounds the unreported Output bytes per FD. + OutputWindow = 256 << 10 +) + +// RelayArgs is the relay's argv. +var RelayArgs = []string{RelayName, "relay"} + +// Exit codes the shim uses when the program did not run or its exit is +// unknown. +const ( + ExitNotFound = 127 + ExitCannotRun = 126 + ExitLost = 255 +) + +// Message types between the shim and the relay. +const ( + TypeRequest uint16 = iota + 1 + TypeAck + TypeSignal + TypeResult +) + +// Message types from the relay to the broker. +const ( + TypeOpen uint16 = iota + 0x11 + TypeInput + TypeInputEnd + TypeWritten + TypeWriteFailed + TypeSignaled + TypeGone +) + +// Message types from the broker to the relay. +const ( + TypeAccept uint16 = iota + 0x21 + TypeStarted + TypeRead + TypeStopInput + TypeOutput + TypeClose + TypeExit + TypeNotice + TypeEnd +) + +// ErrProtocol wraps every IPC violation. +var ErrProtocol = errors.New("processshim: protocol violation") + +// Message is one IPC message. +type Message interface { + messageType() uint16 + encode(*sandboxwire.Encoder) +} + +// RelayMessage is a message from the relay to the broker. +type RelayMessage interface { + Message + // Invocation is the ID of the invocation the message is about. + Invocation() uint64 + fromRelay() +} + +// BrokerMessage is a message from the broker to the relay. +type BrokerMessage interface { + Message + // Invocation is the ID of the invocation the message is about. + Invocation() uint64 + fromBroker() +} + +// Request is the shim's invocation. Every field is what the shim's process +// has; the broker decides what reaches the sandbox. +type Request struct { + Version uint16 + // ExecPath is the path the shim was executed by: AT_EXECFN, else argv[0]. + ExecPath []byte + Argv [][]byte + // Env holds the raw environ entries. + Env [][]byte + Cwd []byte + Umask uint32 +} + +// Ack says the relay owns the three descriptors. +type Ack struct{} + +// Signal reports a signal the shim caught. +type Signal struct{ Number uint16 } + +// Result ends the invocation. A nonzero Signal is the signal that ended the +// remote program, which the shim re-raises; otherwise the shim exits with +// Code. A Message goes to the shim only in a Result sent instead of the Ack, +// and the shim writes it to its stderr as "oac-process-shim: ". +type Result struct { + Signal uint16 + Code uint8 + Message []byte +} + +// WindowSize is a terminal's size. +type WindowSize struct{ Rows, Cols, XPixels, YPixels uint16 } + +// Terminal is the terminal on the shim's fds 0 and 1: its size, and the mode +// saved before any invocation made it raw. +type Terminal struct { + Size WindowSize + Iflag, Oflag, Cflag, Lflag uint32 + // Cc holds the control characters, at most MaxControlChars. + Cc []byte +} + +// Open hands the broker an invocation. +type Open struct { + ID uint64 + Request Request + // Terminal is set when fds 0 and 1 are both terminals. + Terminal *Terminal +} + +// Input is stdin the relay read for one Read. +type Input struct { + ID uint64 + Data []byte +} + +// InputEnd answers a Read at end of file or after a read error. +type InputEnd struct{ ID uint64 } + +// Written reports that Output or Close Seq on FD is done. +type Written struct { + ID uint64 + FD uint8 + Seq uint64 +} + +// WriteFailed reports that writing Output Seq on FD failed with Errno. +type WriteFailed struct { + ID uint64 + FD uint8 + Seq uint64 + Errno uint32 +} + +// Signaled reports a signal from the shim. Size is the terminal's new size +// for SIGWINCH on a terminal. +type Signaled struct { + ID uint64 + Number uint16 + Size *WindowSize +} + +// Gone reports that the shim's connection ended before its Result. +type Gone struct{ ID uint64 } + +// Accept accepts an invocation: the relay acknowledges the shim. +type Accept struct{ ID uint64 } + +// Started says the program runs. +type Started struct{ ID uint64 } + +// Read grants one Input of at most Max bytes, or InputEnd. +type Read struct { + ID uint64 + Max uint32 +} + +// StopInput ends reading stdin. +type StopInput struct{ ID uint64 } + +// Output is data for FD. +type Output struct { + ID uint64 + FD uint8 + Seq uint64 + Data []byte +} + +// Close closes FD after its queued Output. +type Close struct { + ID uint64 + FD uint8 + Seq uint64 +} + +// Mark is the last Output or Close on FD that must be done before the shim +// exits. +type Mark struct { + FD uint8 + Seq uint64 +} + +// Exit ends the shim with Result once Marks are done. +type Exit struct { + ID uint64 + Result Result + Marks []Mark +} + +// Notice is a message for fd 2. +type Notice struct { + ID uint64 + Message []byte +} + +// End closes the invocation. +type End struct{ ID uint64 } + +func (Request) messageType() uint16 { return TypeRequest } +func (Ack) messageType() uint16 { return TypeAck } +func (Signal) messageType() uint16 { return TypeSignal } +func (Result) messageType() uint16 { return TypeResult } +func (Open) messageType() uint16 { return TypeOpen } +func (Input) messageType() uint16 { return TypeInput } +func (InputEnd) messageType() uint16 { return TypeInputEnd } +func (Written) messageType() uint16 { return TypeWritten } +func (WriteFailed) messageType() uint16 { return TypeWriteFailed } +func (Signaled) messageType() uint16 { return TypeSignaled } +func (Gone) messageType() uint16 { return TypeGone } +func (Accept) messageType() uint16 { return TypeAccept } +func (Started) messageType() uint16 { return TypeStarted } +func (Read) messageType() uint16 { return TypeRead } +func (StopInput) messageType() uint16 { return TypeStopInput } +func (Output) messageType() uint16 { return TypeOutput } +func (Close) messageType() uint16 { return TypeClose } +func (Exit) messageType() uint16 { return TypeExit } +func (Notice) messageType() uint16 { return TypeNotice } +func (End) messageType() uint16 { return TypeEnd } + +func (m Open) Invocation() uint64 { return m.ID } +func (m Input) Invocation() uint64 { return m.ID } +func (m InputEnd) Invocation() uint64 { return m.ID } +func (m Written) Invocation() uint64 { return m.ID } +func (m WriteFailed) Invocation() uint64 { return m.ID } +func (m Signaled) Invocation() uint64 { return m.ID } +func (m Gone) Invocation() uint64 { return m.ID } +func (m Accept) Invocation() uint64 { return m.ID } +func (m Started) Invocation() uint64 { return m.ID } +func (m Read) Invocation() uint64 { return m.ID } +func (m StopInput) Invocation() uint64 { return m.ID } +func (m Output) Invocation() uint64 { return m.ID } +func (m Close) Invocation() uint64 { return m.ID } +func (m Exit) Invocation() uint64 { return m.ID } +func (m Notice) Invocation() uint64 { return m.ID } +func (m End) Invocation() uint64 { return m.ID } + +func (Open) fromRelay() {} +func (Input) fromRelay() {} +func (InputEnd) fromRelay() {} +func (Written) fromRelay() {} +func (WriteFailed) fromRelay() {} +func (Signaled) fromRelay() {} +func (Gone) fromRelay() {} +func (Accept) fromBroker() {} +func (Started) fromBroker() {} +func (Read) fromBroker() {} +func (StopInput) fromBroker() {} +func (Output) fromBroker() {} +func (Close) fromBroker() {} +func (Exit) fromBroker() {} +func (Notice) fromBroker() {} +func (End) fromBroker() {} + +func (m Request) encode(e *sandboxwire.Encoder) { + e.U16(m.Version) + e.Bytes(m.ExecPath) + e.Count(len(m.Argv)) + for _, a := range m.Argv { + e.Bytes(a) + } + e.Count(len(m.Env)) + for _, v := range m.Env { + e.Bytes(v) + } + e.Bytes(m.Cwd) + e.U32(m.Umask) +} + +func (Ack) encode(*sandboxwire.Encoder) {} + +func (m Signal) encode(e *sandboxwire.Encoder) { e.U16(m.Number) } + +func (m Result) encode(e *sandboxwire.Encoder) { + e.U16(m.Signal) + e.U8(m.Code) + e.Bytes(m.Message) +} + +func (s WindowSize) encode(e *sandboxwire.Encoder) { + e.U16(s.Rows) + e.U16(s.Cols) + e.U16(s.XPixels) + e.U16(s.YPixels) +} + +func (m Open) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + m.Request.encode(e) + e.Present(m.Terminal != nil) + if t := m.Terminal; t != nil { + t.Size.encode(e) + e.U32(t.Iflag) + e.U32(t.Oflag) + e.U32(t.Cflag) + e.U32(t.Lflag) + e.Bytes(t.Cc) + } +} + +func (m Input) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.Bytes(m.Data) +} + +func (m InputEnd) encode(e *sandboxwire.Encoder) { e.U64(m.ID) } + +func (m Written) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.U8(m.FD) + e.U64(m.Seq) +} + +func (m WriteFailed) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.U8(m.FD) + e.U64(m.Seq) + e.U32(m.Errno) +} + +func (m Signaled) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.U16(m.Number) + e.Present(m.Size != nil) + if m.Size != nil { + m.Size.encode(e) + } +} + +func (m Gone) encode(e *sandboxwire.Encoder) { e.U64(m.ID) } +func (m Accept) encode(e *sandboxwire.Encoder) { e.U64(m.ID) } +func (m Started) encode(e *sandboxwire.Encoder) { e.U64(m.ID) } +func (m StopInput) encode(e *sandboxwire.Encoder) { e.U64(m.ID) } +func (m End) encode(e *sandboxwire.Encoder) { e.U64(m.ID) } + +func (m Read) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.U32(m.Max) +} + +func (m Output) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.U8(m.FD) + e.U64(m.Seq) + e.Bytes(m.Data) +} + +func (m Close) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.U8(m.FD) + e.U64(m.Seq) +} + +func (m Exit) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + m.Result.encode(e) + e.Count(len(m.Marks)) + for _, k := range m.Marks { + e.U8(k.FD) + e.U64(k.Seq) + } +} + +func (m Notice) encode(e *sandboxwire.Encoder) { + e.U64(m.ID) + e.Bytes(m.Message) +} + +// Frame returns m as a frame. +func Frame(m Message) sandboxwire.Frame { + var e sandboxwire.Encoder + m.encode(&e) + return sandboxwire.Frame{Type: m.messageType(), Payload: e.Payload()} +} + +// Decode decodes and validates a frame between the shim and the relay. A +// Request with another Version decodes with only its Version set, so the +// relay can refuse it. +func Decode(f sandboxwire.Frame) (Message, error) { + return decode(f, func(d *sandboxwire.Decoder) (Message, error) { + switch f.Type { + case TypeRequest: + r, err := decodeRequest(d) + if err == nil && r.Version != Version { + return r, nil // the rest is another version's + } + return r, finish(d, err) + case TypeAck: + return Ack{}, d.Finish() + case TypeSignal: + n, err := decodeSignal(d) + return Signal{Number: n}, finish(d, err) + case TypeResult: + r, err := decodeResult(d) + return r, finish(d, err) + } + return nil, errType + }) +} + +// DecodeRelay decodes and validates a frame from the relay. +func DecodeRelay(f sandboxwire.Frame) (RelayMessage, error) { + m, err := decode(f, func(d *sandboxwire.Decoder) (Message, error) { + id, err := decodeID(d) + if err != nil { + return nil, err + } + switch f.Type { + case TypeOpen: + m := Open{ID: id} + m.Request, err = decodeRequest(d) + if err == nil && m.Request.Version != Version { + err = fmt.Errorf("request version %d", m.Request.Version) + } + if err == nil { + m.Terminal, err = decodeTerminal(d) + } + return m, finish(d, err) + case TypeInput: + m := Input{ID: id} + m.Data, err = decodeData(d) + return m, finish(d, err) + case TypeInputEnd: + return InputEnd{ID: id}, d.Finish() + case TypeWritten: + m := Written{ID: id} + m.FD, m.Seq, err = decodeFDSeq(d) + return m, finish(d, err) + case TypeWriteFailed: + m := WriteFailed{ID: id} + m.FD, m.Seq, err = decodeFDSeq(d) + if err == nil { + m.Errno, err = d.U32() + } + if err == nil && m.Errno == 0 { + err = errors.New("errno 0") + } + return m, finish(d, err) + case TypeSignaled: + m := Signaled{ID: id} + m.Number, err = decodeSignal(d) + if err == nil { + m.Size, err = decodeSize(d) + } + return m, finish(d, err) + case TypeGone: + return Gone{ID: id}, d.Finish() + } + return nil, errType + }) + if err != nil { + return nil, err + } + return m.(RelayMessage), nil +} + +// DecodeBroker decodes and validates a frame from the broker. +func DecodeBroker(f sandboxwire.Frame) (BrokerMessage, error) { + m, err := decode(f, func(d *sandboxwire.Decoder) (Message, error) { + id, err := decodeID(d) + if err != nil { + return nil, err + } + switch f.Type { + case TypeAccept: + return Accept{ID: id}, d.Finish() + case TypeStarted: + return Started{ID: id}, d.Finish() + case TypeRead: + m := Read{ID: id} + m.Max, err = d.U32() + if err == nil && (m.Max == 0 || m.Max > sandboxwire.MaxChunk) { + err = fmt.Errorf("read of %d bytes", m.Max) + } + return m, finish(d, err) + case TypeStopInput: + return StopInput{ID: id}, d.Finish() + case TypeOutput: + m := Output{ID: id} + m.FD, m.Seq, err = decodeFDSeq(d) + if err == nil { + m.Data, err = decodeData(d) + } + return m, finish(d, err) + case TypeClose: + m := Close{ID: id} + m.FD, m.Seq, err = decodeFDSeq(d) + return m, finish(d, err) + case TypeExit: + m := Exit{ID: id} + if m.Result, err = decodeResult(d); err == nil { + m.Marks, err = decodeMarks(d) + } + return m, finish(d, err) + case TypeNotice: + m := Notice{ID: id} + m.Message, err = d.Bytes() + if err == nil && (len(m.Message) == 0 || len(m.Message) > MaxMessageBytes) { + err = fmt.Errorf("notice of %d bytes", len(m.Message)) + } + return m, finish(d, err) + case TypeEnd: + return End{ID: id}, d.Finish() + } + return nil, errType + }) + if err != nil { + return nil, err + } + return m.(BrokerMessage), nil +} + +var errType = errors.New("message type") + +func decode(f sandboxwire.Frame, body func(*sandboxwire.Decoder) (Message, error)) (Message, error) { + if f.RequestID != 0 { + return nil, fmt.Errorf("%w: request ID %d", ErrProtocol, f.RequestID) + } + m, err := body(sandboxwire.NewDecoder(f.Payload)) + switch { + case err == errType: + return nil, fmt.Errorf("%w: message type %#x", ErrProtocol, f.Type) + case err != nil: + return nil, fmt.Errorf("%w: %w", ErrProtocol, err) + } + return m, nil +} + +// finish returns err, or the decoder's error for trailing bytes. +func finish(d *sandboxwire.Decoder, err error) error { + if err != nil { + return err + } + return d.Finish() +} + +func decodeID(d *sandboxwire.Decoder) (uint64, error) { + id, err := d.U64() + if err == nil && id == 0 { + err = errors.New("invocation ID 0") + } + return id, err +} + +func decodeRequest(d *sandboxwire.Decoder) (Request, error) { + var r Request + var err error + if r.Version, err = d.U16(); err != nil || r.Version != Version { + return Request{Version: r.Version}, err + } + if r.ExecPath, err = d.Bytes(); err != nil { + return r, err + } + if r.Argv, err = decodeList(d); err != nil { + return r, err + } + if r.Env, err = decodeList(d); err != nil { + return r, err + } + if r.Cwd, err = d.Bytes(); err != nil { + return r, err + } + if r.Umask, err = d.U32(); err != nil { + return r, err + } + switch { + case slices.Contains(r.ExecPath, 0) || slices.Contains(r.Cwd, 0): + return r, errors.New("NUL in path") + case len(r.Cwd) == 0 || r.Cwd[0] != '/': + return r, errors.New("cwd is not absolute") + case r.Umask > 0o777: + return r, fmt.Errorf("umask %#o", r.Umask) + } + for _, b := range slices.Concat(r.Argv, r.Env) { + if slices.Contains(b, 0) { + return r, errors.New("NUL in argv or environment") + } + } + return r, nil +} + +func decodeList(d *sandboxwire.Decoder) ([][]byte, error) { + n, err := d.Count(MaxFrameBytes / 4) + if err != nil { + return nil, err + } + list := make([][]byte, n) + for i := range list { + if list[i], err = d.Bytes(); err != nil { + return nil, err + } + } + return list, nil +} + +func decodeSignal(d *sandboxwire.Decoder) (uint16, error) { + n, err := d.U16() + if err == nil && (n == 0 || n > 64) { + err = fmt.Errorf("signal %d", n) + } + return n, err +} + +func decodeResult(d *sandboxwire.Decoder) (Result, error) { + var r Result + var err error + if r.Signal, err = d.U16(); err != nil { + return r, err + } + if r.Code, err = d.U8(); err != nil { + return r, err + } + if r.Message, err = d.Bytes(); err != nil { + return r, err + } + switch { + case r.Signal > 64 || (r.Signal != 0 && r.Code != 0): + return r, fmt.Errorf("signal %d with code %d", r.Signal, r.Code) + case len(r.Message) > MaxMessageBytes: + return r, fmt.Errorf("message of %d bytes", len(r.Message)) + } + return r, nil +} + +func decodeSize(d *sandboxwire.Decoder) (*WindowSize, error) { + ok, err := d.Present() + if err != nil || !ok { + return nil, err + } + var s WindowSize + for _, v := range []*uint16{&s.Rows, &s.Cols, &s.XPixels, &s.YPixels} { + if *v, err = d.U16(); err != nil { + return nil, err + } + } + return &s, nil +} + +func decodeTerminal(d *sandboxwire.Decoder) (*Terminal, error) { + size, err := decodeSize(d) + if err != nil || size == nil { + return nil, err + } + t := &Terminal{Size: *size} + for _, v := range []*uint32{&t.Iflag, &t.Oflag, &t.Cflag, &t.Lflag} { + if *v, err = d.U32(); err != nil { + return nil, err + } + } + if t.Cc, err = d.Bytes(); err != nil { + return nil, err + } + if len(t.Cc) > MaxControlChars { + return nil, fmt.Errorf("%d control characters", len(t.Cc)) + } + return t, nil +} + +func decodeData(d *sandboxwire.Decoder) ([]byte, error) { + b, err := d.Bytes() + if err == nil && (len(b) == 0 || len(b) > sandboxwire.MaxChunk) { + err = fmt.Errorf("data of %d bytes", len(b)) + } + return b, err +} + +func decodeFD(d *sandboxwire.Decoder) (uint8, error) { + fd, err := d.U8() + if err == nil && fd != 1 && fd != 2 { + err = fmt.Errorf("fd %d", fd) + } + return fd, err +} + +func decodeFDSeq(d *sandboxwire.Decoder) (uint8, uint64, error) { + fd, err := decodeFD(d) + if err != nil { + return 0, 0, err + } + seq, err := d.U64() + if err == nil && seq == 0 { + err = errors.New("sequence 0") + } + return fd, seq, err +} + +func decodeMarks(d *sandboxwire.Decoder) ([]Mark, error) { + n, err := d.Count(2) + if err != nil || n == 0 { + return nil, err + } + marks := make([]Mark, n) + for i := range marks { + if marks[i].FD, marks[i].Seq, err = decodeFDSeq(d); err != nil { + return nil, err + } + if i > 0 && marks[i].FD == marks[0].FD { + return nil, fmt.Errorf("two marks on fd %d", marks[i].FD) + } + } + return marks, nil +} diff --git a/apps/daemon/internal/processshim/ipc_test.go b/apps/daemon/internal/processshim/ipc_test.go new file mode 100644 index 00000000..97392f5d --- /dev/null +++ b/apps/daemon/internal/processshim/ipc_test.go @@ -0,0 +1,152 @@ +package processshim + +import ( + "bytes" + "encoding/binary" + "encoding/hex" + "errors" + "os" + "path/filepath" + "reflect" + "strings" + "testing" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +func b(s string) []byte { return []byte(s) } + +var fixtures = []struct { + name string + msg Message +}{ + {"request.hex", Request{ + Version: Version, + ExecPath: b("/.oac/bin/git"), + Argv: [][]byte{b("git"), b("status")}, + Env: [][]byte{b("HOME=/home/u")}, + Cwd: b("/work"), + Umask: 0o022, + }}, + {"ack.hex", Ack{}}, + {"signal.hex", Signal{Number: 2}}, + {"result_code.hex", Result{Code: 3, Message: []byte{}}}, + {"result_signal.hex", Result{Signal: 15, Message: []byte{}}}, + {"result_refused.hex", Result{Code: ExitNotFound, Message: b("not found")}}, + {"open.hex", Open{ID: 1, Request: Request{ + Version: Version, + ExecPath: b("/.oac/bin/sh"), + Argv: [][]byte{b("sh")}, + Env: [][]byte{}, + Cwd: b("/w"), + Umask: 0o022, + }, Terminal: &Terminal{ + Size: WindowSize{Rows: 24, Cols: 80}, + Iflag: 0x500, Oflag: 0x5, Cflag: 0xbf, Lflag: 0x8a3b, + Cc: []byte{3, 28}, + }}}, + {"input.hex", Input{ID: 1, Data: b("hi\n")}}, + {"input_end.hex", InputEnd{ID: 1}}, + {"written.hex", Written{ID: 1, FD: 1, Seq: 7}}, + {"write_failed.hex", WriteFailed{ID: 1, FD: 1, Seq: 8, Errno: 13}}, + {"signaled.hex", Signaled{ID: 1, Number: 28, Size: &WindowSize{Rows: 30, Cols: 100}}}, + {"gone.hex", Gone{ID: 2}}, + {"accept.hex", Accept{ID: 1}}, + {"started.hex", Started{ID: 1}}, + {"read.hex", Read{ID: 1, Max: 64 << 10}}, + {"stop_input.hex", StopInput{ID: 1}}, + {"output.hex", Output{ID: 1, FD: 2, Seq: 5, Data: b("err\n")}}, + {"close.hex", Close{ID: 1, FD: 1, Seq: 9}}, + {"exit.hex", Exit{ID: 1, Result: Result{Code: 3, Message: []byte{}}, Marks: []Mark{{FD: 1, Seq: 9}, {FD: 2, Seq: 6}}}}, + {"notice.hex", Notice{ID: 1, Message: b("link lost")}}, + {"end.hex", End{ID: 1}}, +} + +// decoders are the three decoders: shim and relay, relay to broker, and +// broker to relay. +var decoders = []func(sandboxwire.Frame) (Message, error){ + Decode, + func(f sandboxwire.Frame) (Message, error) { return DecodeRelay(f) }, + func(f sandboxwire.Frame) (Message, error) { return DecodeBroker(f) }, +} + +// decoderFor picks the decoder of a message type. +func decoderFor(t uint16) func(sandboxwire.Frame) (Message, error) { + return decoders[min(t>>4, 2)] +} + +func readHexFixture(t testing.TB, name string) []byte { + t.Helper() + raw, err := os.ReadFile(filepath.Join("testdata", name)) + if err != nil { + t.Fatal(err) + } + var digits strings.Builder + for _, line := range strings.Split(string(raw), "\n") { + line, _, _ = strings.Cut(line, "#") + digits.WriteString(strings.Join(strings.Fields(line), "")) + } + frame, err := hex.DecodeString(digits.String()) + if err != nil { + t.Fatal(err) + } + return frame +} + +func TestGoldenFixtures(t *testing.T) { + for _, fx := range fixtures { + t.Run(fx.name, func(t *testing.T) { + want := readHexFixture(t, fx.name) + var got bytes.Buffer + if err := sandboxwire.WriteFrame(&got, Frame(fx.msg)); err != nil { + t.Fatal(err) + } + if !bytes.Equal(got.Bytes(), want) { + t.Fatalf("encoded\n got %x\nwant %x", got.Bytes(), want) + } + f, err := sandboxwire.ReadFrame(bytes.NewReader(want), MaxFrameBytes) + if err != nil { + t.Fatal(err) + } + m, err := decoderFor(f.Type)(f) + if err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(m, fx.msg) { + t.Fatalf("decoded\n got %#v\nwant %#v", m, fx.msg) + } + }) + } +} + +// FuzzDecode takes a message type and a payload and runs every decoder on +// them. Whatever decodes must re-encode to the same bytes, except a Request +// of another version, which decodes to its version alone. +func FuzzDecode(f *testing.F) { + for _, fx := range fixtures { + fr := Frame(fx.msg) + f.Add(binary.BigEndian.AppendUint16(nil, fr.Type), fr.Payload) + } + f.Add([]byte{0, 1}, []byte{0, 2}) + f.Fuzz(func(t *testing.T, tag, payload []byte) { + if len(tag) != 2 { + return + } + typ := binary.BigEndian.Uint16(tag) + for _, decode := range decoders { + m, err := decode(sandboxwire.Frame{Type: typ, Payload: payload}) + if err != nil { + if !errors.Is(err, ErrProtocol) { + t.Fatalf("error %v does not wrap ErrProtocol", err) + } + continue + } + if r, ok := m.(Request); ok && r.Version != Version { + continue + } + if got := Frame(m); got.Type != typ || !bytes.Equal(got.Payload, payload) { + t.Fatalf("%T re-encoded as %#x %x, decoded from %#x %x", m, got.Type, got.Payload, typ, payload) + } + } + }) +} diff --git a/apps/daemon/internal/processshim/layout_test.go b/apps/daemon/internal/processshim/layout_test.go new file mode 100644 index 00000000..bef77865 --- /dev/null +++ b/apps/daemon/internal/processshim/layout_test.go @@ -0,0 +1,23 @@ +package processshim + +import ( + "path" + "testing" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/agent" +) + +func TestSocketPathFollowsTheViewLayout(t *testing.T) { + if want := path.Join(agent.ViewPrivateRoot, agent.ViewRunName, SocketName); SocketPath != want { + t.Fatalf("SocketPath = %q, want %q", SocketPath, want) + } +} + +func TestRelayPathFollowsTheViewLayout(t *testing.T) { + if RelayName != agent.ViewRelayName { + t.Fatalf("RelayName = %q, want %q", RelayName, agent.ViewRelayName) + } + if want := path.Join(agent.ViewPrivateRoot, agent.ViewShimName, RelayName); RelayPath != want { + t.Fatalf("RelayPath = %q, want %q", RelayPath, want) + } +} diff --git a/apps/daemon/internal/processshim/relay_linux.go b/apps/daemon/internal/processshim/relay_linux.go new file mode 100644 index 00000000..50094386 --- /dev/null +++ b/apps/daemon/internal/processshim/relay_linux.go @@ -0,0 +1,706 @@ +//go:build linux + +package processshim + +import ( + "bufio" + "errors" + "fmt" + "net" + "os" + "os/signal" + "slices" + "sync" + "time" + + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// Relaying reports whether the process runs as the relay: sessionview +// executed RelayPath with argv RelayArgs. It must run before Relay. +func Relaying() bool { + return slices.Equal(os.Args, RelayArgs) && string(atExecFn()) == RelayPath +} + +const ( + // handshakeTimeout bounds reading a shim's Request. + handshakeTimeout = 10 * time.Second + // relayWait bounds how long the relay, once the broker is gone, keeps + // answering waiting shims before it exits. + relayWait = time.Second +) + +// Relay runs the Session's process relay with the broker's connection at +// RelayBrokerFD and the listening socket at RelayListenerFD, until the +// broker's connection ends. It returns the process's exit code. +func Relay() int { + // Other processes of the Session user may not trace the relay or open + // its descriptors through /proc. + unix.Prctl(unix.PR_SET_DUMPABLE, 0, 0, 0, 0) + // The relay serves until the broker goes; signals the view forwards to + // every process are not for it. + signal.Ignore() + r, err := newRelay() + if err != nil { + return 1 + } + go r.accept() + r.serveBroker() + r.lose() + return 0 +} + +// relay serves the Session's shims. +type relay struct { + broker *net.UnixConn + ln *net.UnixListener + // sendMu orders the frames to the broker. An invocation's ID is + // allocated and its Open written under it in one step, so Open IDs reach + // the broker in increasing order. + sendMu sync.Mutex + terms terminals + + mu sync.Mutex + invs map[uint64]*invocation + lastID uint64 + gone bool // the broker's connection ended + live sync.WaitGroup // the invocations' control goroutines + + // Test seams: publishing runs between an invocation's ID and its Open, + // gating before an output pump waits for the terminal to be raw, and + // makingRaw before the control goroutine makes a terminal raw. + publishing func(Request) + gating func() + makingRaw func() +} + +func newRelay() (*relay, error) { + bf := os.NewFile(RelayBrokerFD, "broker") + c, err := net.FileConn(bf) + bf.Close() + if err != nil { + return nil, err + } + uc, ok := c.(*net.UnixConn) + if !ok { + c.Close() + return nil, errors.New("the broker's connection is not a Unix socket") + } + lf := os.NewFile(RelayListenerFD, "listener") + l, err := net.FileListener(lf) + lf.Close() + if err != nil { + uc.Close() + return nil, err + } + ul, ok := l.(*net.UnixListener) + if !ok { + uc.Close() + l.Close() + return nil, errors.New("the listener is not a Unix socket") + } + ul.SetUnlinkOnClose(false) + return &relay{broker: uc, ln: ul, invs: map[uint64]*invocation{}}, nil +} + +func (r *relay) send(m RelayMessage) error { + r.sendMu.Lock() + defer r.sendMu.Unlock() + return sandboxwire.WriteFrame(r.broker, Frame(m)) +} + +func (r *relay) accept() { + for { + c, err := r.ln.AcceptUnix() + if err != nil { + if errors.Is(err, net.ErrClosed) { + return + } + time.Sleep(10 * time.Millisecond) // out of descriptors, for example + continue + } + go r.handshake(c) + } +} + +// handshake reads a shim's Request and hands the invocation to the broker, +// or refuses it. +func (r *relay) handshake(c *net.UnixConn) { + conn := NewConn(c) + c.SetDeadline(time.Now().Add(handshakeTimeout)) + req, fds, err := conn.ReadRequest() // closes the descriptors on error + if err != nil { + conn.Close() + return + } + inv, refusal := r.open(conn, req, fds) + if inv == nil { + closeAll(fds[:]) + msg := []byte(refusal) + conn.Send(Result{Code: ExitCannotRun, Message: msg[:min(len(msg), MaxMessageBytes)]}) + conn.Close() + return + } + c.SetDeadline(time.Time{}) + inv.start() +} + +// open prepares, registers and publishes an invocation, or returns why it is +// refused. On refusal the caller still owns fds. +func (r *relay) open(conn *Conn, req Request, fds [3]int) (*invocation, string) { + if req.Version != Version { + return nil, fmt.Sprintf("IPC version %d is not %d", req.Version, Version) + } + inv := &invocation{r: r, conn: conn, raw: make(chan struct{}), wake: make(chan struct{}, 1)} + var err error + if inv.end, err = newStopFlag(); err != nil { + return nil, err.Error() + } + in := &input{inv: inv, credit: make(chan uint32, 1)} + if in.stop, err = newStopFlag(); err != nil { + inv.end.close() + return nil, err.Error() + } + inv.in = in + var opened []endpoint + refuse := func(msg string) (*invocation, string) { + for _, ep := range opened { + if ep.io >= 0 && ep.io != ep.fd { + unix.Close(ep.io) + } + } + inv.term.close() + inv.end.close() + in.stop.close() + return nil, msg + } + for i, fd := range fds { + ep, err := openEndpoint(fd, i > 0) + if err != nil { + return refuse(fmt.Sprintf("descriptor %d: %v", i, err)) + } + opened = append(opened, ep) + if i == 0 { + in.ep = ep + continue + } + inv.out[i] = &output{inv: inv, fd: uint8(i), ep: ep, wake: make(chan struct{}, 1), changed: make(chan struct{})} + } + if inv.term, err = openTerminal(&r.terms, fds[0], fds[1]); err != nil { + return refuse(fmt.Sprintf("terminal: %v", err)) + } + var mode *Terminal + if inv.term != nil { + mode = inv.term.mode() + } else { + inv.rawOnce.Do(func() { close(inv.raw) }) + } + r.sendMu.Lock() + defer r.sendMu.Unlock() + r.mu.Lock() + switch { + case r.gone: + r.mu.Unlock() + return refuse("the process broker is unavailable") + case len(r.invs) >= MaxInvocations: + r.mu.Unlock() + return refuse(fmt.Sprintf("more than %d programs are running", MaxInvocations)) + } + r.lastID++ + inv.id = r.lastID + r.invs[inv.id] = inv + r.live.Add(1) + r.mu.Unlock() + if r.publishing != nil { + r.publishing(req) + } + // The broker's messages for the ID queue until start runs the + // invocation. A failed write means the broker is gone, which ends it. + sandboxwire.WriteFrame(r.broker, Frame(Open{ID: inv.id, Request: req, Terminal: mode})) + return inv, "" +} + +// serveBroker dispatches the broker's messages until its connection ends. +// It never blocks on an invocation. +func (r *relay) serveBroker() { + br := bufio.NewReaderSize(r.broker, 64<<10) + for { + f, err := sandboxwire.ReadFrame(br, MaxFrameBytes) + if err != nil { + return + } + m, err := DecodeBroker(f) + if err != nil { + return + } + r.mu.Lock() + inv := r.invs[m.Invocation()] + r.mu.Unlock() + if inv != nil { + inv.receive(m) + } + } +} + +// lose ends every invocation after the broker's connection ended, and +// waits a bounded time for them. +func (r *relay) lose() { + r.ln.Close() + r.mu.Lock() + r.gone = true + invs := make([]*invocation, 0, len(r.invs)) + for _, inv := range r.invs { + invs = append(invs, inv) + } + r.mu.Unlock() + for _, inv := range invs { + inv.end.set() + inv.control(lost{}) + } + done := make(chan struct{}) + go func() { + r.live.Wait() + close(done) + }() + select { + case <-done: + case <-time.After(relayWait): + } +} + +// invocation is one shim's invocation in the relay. +type invocation struct { + r *relay + id uint64 + conn *Conn + term *terminal // nil for pipes + in *input + out [3]*output // 1 and 2 + // end is set by End or the broker's loss; it stops every pump and wait. + end *stopFlag + // raw closes once stdin may be read and output written: the terminal + // is raw, or there is none. + raw chan struct{} + rawOnce sync.Once + pumps sync.WaitGroup // the output pumps + + ctlMu sync.Mutex + ctl []any + wake chan struct{} + + mu sync.Mutex + acked bool + answered bool // the shim has its Result or is gone +} + +// lost is the control item for the broker's loss. +type lost struct{} + +// start runs the published invocation. The output pumps are counted before +// run starts, as its finish waits for them. +func (inv *invocation) start() { + inv.pumps.Add(2) + go inv.run() + go inv.out[1].run() + go inv.out[2].run() + go inv.in.run() + go inv.readShim() +} + +// receive takes one broker message without blocking. +func (inv *invocation) receive(m BrokerMessage) { + switch m := m.(type) { + case Output: + inv.out[m.FD].push(item{seq: m.Seq, data: m.Data}) + case Close: + inv.out[m.FD].push(item{seq: m.Seq, close: true}) + case Read: + select { + case inv.in.credit <- m.Max: + default: + } + case StopInput: + inv.in.stop.set() + case Exit: + inv.in.stop.set() + inv.control(m) + case End: + inv.end.set() + inv.control(m) + default: + inv.control(m) + } +} + +func (inv *invocation) control(m any) { + inv.ctlMu.Lock() + inv.ctl = append(inv.ctl, m) + inv.ctlMu.Unlock() + select { + case inv.wake <- struct{}{}: + default: + } +} + +func (inv *invocation) nextControl() any { + for { + inv.ctlMu.Lock() + if len(inv.ctl) > 0 { + m := inv.ctl[0] + inv.ctl[0] = nil + inv.ctl = inv.ctl[1:] + inv.ctlMu.Unlock() + return m + } + inv.ctlMu.Unlock() + <-inv.wake + } +} + +// run handles the control messages in order until End or the broker's loss. +func (inv *invocation) run() { + defer inv.r.live.Done() + for { + switch m := inv.nextControl().(type) { + case Accept: + inv.mu.Lock() + inv.acked = true + inv.mu.Unlock() + inv.conn.Send(Ack{}) // a failure means the shim is gone; readShim reports it + case Started: + if inv.term != nil { + if inv.r.makingRaw != nil { + inv.r.makingRaw() + } + inv.term.makeRaw() // a terminal that stays cooked still works + } + inv.rawOnce.Do(func() { close(inv.raw) }) + case Exit: + inv.exit(m) + case Notice: + inv.out[2].message(m.Message) + case End: + inv.finish("") + return + case lost: + inv.finish("the process broker stopped") + return + } + } +} + +// exit answers the shim once the output before the exit is written. When +// End cuts that wait short, the shim gets ExitLost. +func (inv *invocation) exit(m Exit) { + complete := inv.waitMarks(m.Marks) + inv.term.restore() + res := m.Result + if !complete { + res = Result{Code: ExitLost} + } + inv.mu.Lock() + acked := inv.acked + inv.mu.Unlock() + if acked && len(res.Message) > 0 { + inv.out[2].message(res.Message) + res.Message = nil + } + inv.answer(res) +} + +func (inv *invocation) waitMarks(marks []Mark) bool { + for _, m := range marks { + o := inv.out[m.FD] + for { + done, changed := o.reached(m.Seq) + if done { + break + } + select { + case <-changed: + case <-inv.end.c: + if done, _ := o.reached(m.Seq); !done { + return false + } + } + } + } + return true +} + +// answer sends the shim its one Result, unless it has one or is gone. +func (inv *invocation) answer(r Result) { + inv.mu.Lock() + if inv.answered { + inv.mu.Unlock() + return + } + inv.answered = true + inv.mu.Unlock() + inv.conn.Send(r) + inv.conn.Close() +} + +// finish ends the invocation: it stops the pumps, answers a waiting shim +// with ExitLost, restores the terminal and closes every descriptor. +func (inv *invocation) finish(reason string) { + inv.end.set() + inv.in.stop.set() + if reason != "" { + inv.out[2].message([]byte(reason)) + } + inv.answer(Result{Code: ExitLost}) + inv.term.close() + inv.rawOnce.Do(func() { close(inv.raw) }) + inv.r.mu.Lock() + delete(inv.r.invs, inv.id) + inv.r.mu.Unlock() + inv.pumps.Wait() + inv.out[1].closeEP() + inv.out[2].closeEP() + inv.end.close() +} + +// readShim forwards the shim's signals and reports its loss. +func (inv *invocation) readShim() { + for { + m, err := inv.conn.ReadMessage() + s, ok := m.(Signal) + if err != nil || !ok { + inv.shimGone() + return + } + sig := Signaled{ID: inv.id, Number: s.Number} + if inv.term != nil && s.Number == uint16(unix.SIGWINCH) { + size := inv.term.size() + sig.Size = &size + } + inv.r.send(sig) + } +} + +// shimGone handles the end of the shim's connection before its Result. +func (inv *invocation) shimGone() { + inv.mu.Lock() + if inv.answered { + inv.mu.Unlock() + return + } + inv.answered = true + inv.mu.Unlock() + inv.conn.Close() + inv.in.stop.set() + inv.term.restore() + inv.r.send(Gone{ID: inv.id}) +} + +// input reads descriptor 0 for the broker, once per Read. It owns the +// descriptor and closes it when it returns. +type input struct { + inv *invocation + ep endpoint + credit chan uint32 + // stop is set by StopInput, Exit, the shim's loss and the end. + stop *stopFlag +} + +func (in *input) run() { + defer in.stop.close() + defer in.ep.close() + select { + case <-in.inv.raw: + case <-in.stop.c: + return + } + buf := make([]byte, sandboxwire.MaxChunk) + for { + var n uint32 + select { + case n = <-in.credit: + case <-in.stop.c: + return + } + got, err := readFD(in.ep, buf[:n], in.stop) + switch { + case err == errStopped: + return + case got == 0 || err != nil: + in.inv.r.send(InputEnd{ID: in.inv.id}) + return + } + if in.inv.r.send(Input{ID: in.inv.id, Data: buf[:got]}) != nil { + return + } + } +} + +// item is an Output or Close for an output pump. A local Close closes the +// descriptor without a report. +type item struct { + seq uint64 + data []byte + close bool + local bool +} + +// output writes one descriptor's Output in order and reports each write. +// Only its pump closes the descriptor before the invocation finishes. +type output struct { + inv *invocation + fd uint8 + ep endpoint + + epMu sync.Mutex + closed bool + + mu sync.Mutex + queue []item + wake chan struct{} + done uint64 // the last Seq written or closed + broken bool // a write failed; nothing more is reported + changed chan struct{} +} + +func (o *output) push(it item) { + if o == nil { + return + } + o.mu.Lock() + o.queue = append(o.queue, it) + o.mu.Unlock() + select { + case o.wake <- struct{}{}: + default: + } +} + +func (o *output) next() (item, bool) { + for { + o.mu.Lock() + if len(o.queue) > 0 { + it := o.queue[0] + o.queue[0] = item{} + o.queue = o.queue[1:] + o.mu.Unlock() + return it, true + } + o.mu.Unlock() + select { + case <-o.wake: + case <-o.inv.end.c: + return item{}, false + } + } +} + +func (o *output) run() { + defer o.inv.pumps.Done() + id := o.inv.id + for { + it, ok := o.next() + if !ok { + return + } + if it.close { + o.closeEP() + if it.local { + continue + } + if o.fd == 1 && o.inv.term != nil { + o.inv.out[2].push(item{close: true, local: true}) // the merged stream never uses it + } + if o.progress(it.seq, false) { + o.inv.r.send(Written{ID: id, FD: o.fd, Seq: it.seq}) + } + continue + } + if o.isBroken() || o.isClosed() { + continue + } + // Until the terminal is raw, its output processing would translate + // the remote terminal's output a second time. + if o.inv.r.gating != nil { + o.inv.r.gating() + } + select { + case <-o.inv.raw: + case <-o.inv.end.c: + return + } + err := writeFD(o.ep, it.data, o.inv.end) + switch { + case err == errStopped: + return + case err != nil: + if o.progress(it.seq, true) { + o.inv.r.send(WriteFailed{ID: id, FD: o.fd, Seq: it.seq, Errno: errnoOf(err)}) + } + case o.progress(it.seq, false): + o.inv.r.send(Written{ID: id, FD: o.fd, Seq: it.seq}) + } + } +} + +// progress records seq as done, or the descriptor as broken, and reports +// whether the broker is told. +func (o *output) progress(seq uint64, failed bool) bool { + o.mu.Lock() + defer o.mu.Unlock() + if o.broken { + return false + } + o.broken = failed + o.done = max(o.done, seq) + close(o.changed) + o.changed = make(chan struct{}) + return true +} + +// reached reports whether Output and Close through seq are done or the +// descriptor is broken, and a channel that closes on the next progress. +func (o *output) reached(seq uint64) (bool, <-chan struct{}) { + o.mu.Lock() + defer o.mu.Unlock() + return o.broken || o.done >= seq, o.changed +} + +func (o *output) isBroken() bool { + o.mu.Lock() + defer o.mu.Unlock() + return o.broken +} + +func (o *output) isClosed() bool { + o.epMu.Lock() + defer o.epMu.Unlock() + return o.closed +} + +func (o *output) closeEP() { + o.epMu.Lock() + defer o.epMu.Unlock() + if !o.closed { + o.closed = true + o.ep.close() + } +} + +// message writes "oac-process-shim: msg" as far as the descriptor takes it +// now. A failed message does not break the descriptor. +func (o *output) message(msg []byte) { + o.epMu.Lock() + defer o.epMu.Unlock() + if !o.closed { + tryWrite(o.ep, fmt.Appendf(nil, "oac-process-shim: %s\n", msg)) + } +} + +func errnoOf(err error) uint32 { + var errno unix.Errno + if errors.As(err, &errno) && errno != 0 { + return uint32(errno) + } + return uint32(unix.EIO) +} diff --git a/apps/daemon/internal/processshim/relay_linux_test.go b/apps/daemon/internal/processshim/relay_linux_test.go new file mode 100644 index 00000000..16288c00 --- /dev/null +++ b/apps/daemon/internal/processshim/relay_linux_test.go @@ -0,0 +1,327 @@ +//go:build linux + +package processshim + +import ( + "bufio" + "fmt" + "net" + "os" + "path/filepath" + "runtime" + "strings" + "sync" + "testing" + "time" + + "github.com/creack/pty" + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +// testRelay runs a relay in this process with its broker's end at broker. +type testRelay struct { + sock string + broker *net.UnixConn + in *bufio.Reader +} + +// startTestRelay starts a relay that configure may give test seams. +func startTestRelay(t *testing.T, configure func(*relay)) *testRelay { + t.Helper() + sv, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + if err != nil { + t.Fatal(err) + } + ends := [2]*net.UnixConn{} + for i, fd := range sv { + f := os.NewFile(uintptr(fd), "broker") + c, err := net.FileConn(f) + f.Close() + if err != nil { + t.Fatal(err) + } + ends[i] = c.(*net.UnixConn) + } + sock := filepath.Join(t.TempDir(), SocketName) + ln, err := net.ListenUnix("unix", &net.UnixAddr{Name: sock, Net: "unix"}) + if err != nil { + t.Fatal(err) + } + r := &relay{broker: ends[1], ln: ln, invs: map[uint64]*invocation{}} + if configure != nil { + configure(r) + } + done := make(chan struct{}) + go r.accept() + go func() { + defer close(done) + r.serveBroker() + r.lose() + r.broker.Close() + }() + t.Cleanup(func() { + ends[0].Close() + select { + case <-done: + case <-time.After(10 * time.Second): + t.Error("the relay still runs after its broker left") + } + }) + return &testRelay{sock: sock, broker: ends[0], in: bufio.NewReader(ends[0])} +} + +// read returns the relay's next message to the broker. +func (r *testRelay) read() (RelayMessage, error) { + f, err := sandboxwire.ReadFrame(r.in, MaxFrameBytes) + if err != nil { + return nil, err + } + return DecodeRelay(f) +} + +func (r *testRelay) send(ms ...BrokerMessage) error { + for _, m := range ms { + if err := sandboxwire.WriteFrame(r.broker, Frame(m)); err != nil { + return err + } + } + return nil +} + +// shim hands the relay an invocation of name with fds, as the shim does. +func (r *testRelay) shim(t *testing.T, name string, fds [3]int) *Conn { + t.Helper() + c, err := net.DialUnix("unix", nil, &net.UnixAddr{Name: r.sock, Net: "unix"}) + if err != nil { + t.Fatal(err) + } + c.SetDeadline(time.Now().Add(10 * time.Second)) + conn := NewConn(c) + t.Cleanup(func() { conn.Close() }) + req := Request{Version: Version, ExecPath: []byte("/.oac/bin/" + name), Argv: [][]byte{[]byte(name)}, Env: [][]byte{}, Cwd: []byte("/"), Umask: 0o022} + if err := conn.SendRequest(req, fds); err != nil { + t.Fatal(err) + } + return conn +} + +// finished reads the shim's Ack and Result. +func finished(conn *Conn) (Result, error) { + m, err := conn.ReadMessage() + if err != nil { + return Result{}, err + } + if _, ok := m.(Ack); !ok { + return Result{}, fmt.Errorf("got %#v before the Ack", m) + } + m, err = conn.ReadMessage() + if err != nil { + return Result{}, err + } + res, ok := m.(Result) + if !ok { + return Result{}, fmt.Errorf("got %#v, want a Result", m) + } + return res, nil +} + +func devNull(t *testing.T) [3]int { + t.Helper() + f, err := os.OpenFile(os.DevNull, os.O_RDWR, 0) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { f.Close() }) + fd := int(f.Fd()) + return [3]int{fd, fd, fd} +} + +// While one invocation is between its ID and its Open, another handshake +// waits for both steps, so the broker sees the Opens with increasing IDs, as +// it requires. +func TestOpensReachTheBrokerInIDOrder(t *testing.T) { + paused, release := make(chan struct{}), make(chan struct{}) + r := startTestRelay(t, func(r *relay) { + r.publishing = func(req Request) { + if string(req.Argv[0]) == "first" { + close(paused) + <-release + } + } + }) + broken := make(chan error, 1) + secondOpened := make(chan struct{}) + go func() { + var last uint64 + for { + m, err := r.read() + if err != nil { + return + } + open, ok := m.(Open) + if !ok { + continue + } + if string(open.Request.Argv[0]) == "second" { + close(secondOpened) + } + if open.ID <= last { + // The broker ends the Session at an Open ID that does not + // increase. + broken <- fmt.Errorf("Open %d after %d", open.ID, last) + r.broker.Close() + return + } + last = open.ID + id := open.ID + r.send(Accept{ID: id}, Started{ID: id}, Exit{ID: id, Result: Result{Code: 0}}, End{ID: id}) + } + }() + first := r.shim(t, "first", devNull(t)) + select { + case <-paused: + case <-time.After(10 * time.Second): + t.Fatal("the first handshake never reached publication") + } + // The second handshake either waits for the first's publication or, if + // the steps were apart, publishes its own Open first. + second := r.shim(t, "second", devNull(t)) + settled := func() bool { + select { + case <-secondOpened: + return true + default: + return waitsForLock("processshim.(*relay).open(") + } + } + for deadline := time.Now().Add(10 * time.Second); !settled(); time.Sleep(time.Millisecond) { + if time.Now().After(deadline) { + t.Fatal("the second handshake neither waited nor published") + } + } + close(release) + for name, conn := range map[string]*Conn{"first": first, "second": second} { + if res, err := finished(conn); err != nil || res.Code != 0 { + t.Errorf("%s: Result %+v, %v", name, res, err) + } + } + select { + case err := <-broken: + t.Fatal(err) + default: + } +} + +// waitsForLock reports whether a goroutine in fn waits to lock a mutex. +func waitsForLock(fn string) bool { + buf := make([]byte, 1<<16) + for { + n := runtime.Stack(buf, true) + if n < len(buf) { + buf = buf[:n] + break + } + buf = make([]byte, 2*len(buf)) + } + for _, g := range strings.Split(string(buf), "\n\n") { + if strings.Contains(g, "[sync.Mutex.Lock") && strings.Contains(g, fn) { + return true + } + } + return false +} + +// Output for a terminal is written only once the terminal is raw, so the +// local terminal never processes the remote terminal's output again, even +// when the output reaches its pump before the control goroutine takes +// Started. +func TestTerminalOutputWaitsForRawMode(t *testing.T) { + ptm, pts, err := pty.Open() + if err != nil { + t.Fatal(err) + } + defer ptm.Close() + defer pts.Close() + fd := int(pts.Fd()) + tio, err := unix.IoctlGetTermios(fd, unix.TCGETS) + if err != nil { + t.Fatal(err) + } + tio.Oflag |= unix.OPOST | unix.ONLCR // cooked output turns \n into \r\n + if err := unix.IoctlSetTermios(fd, unix.TCSETS, tio); err != nil { + t.Fatal(err) + } + gated, written := make(chan struct{}), make(chan struct{}) + var gateOnce sync.Once + r := startTestRelay(t, func(r *relay) { + r.gating = func() { gateOnce.Do(func() { close(gated) }) } + // Hold the control goroutine in Started until the output pump holds + // the output, or a pump that does not wait for raw mode wrote it. + r.makingRaw = func() { + select { + case <-gated: + case <-written: + case <-time.After(10 * time.Second): + t.Error("the output never reached its pump") + } + select { + case <-written: + t.Error("the output was written before the terminal was raw") + default: + } + } + }) + const out = "a\r\nb\r\n" // the remote terminal's output + go func() { + for { + m, err := r.read() + if err != nil { + return + } + switch m := m.(type) { + case Open: + id := m.ID + r.send(Accept{ID: id}, Started{ID: id}, Output{ID: id, FD: 1, Seq: 1, Data: []byte(out)}) + case Written: + close(written) + r.send(Exit{ID: m.ID, Result: Result{Code: 0}, Marks: []Mark{{FD: 1, Seq: 1}}}, End{ID: m.ID}) + } + } + }() + conn := r.shim(t, "sh", [3]int{fd, fd, fd}) + if res, err := finished(conn); err != nil || res.Code != 0 { + t.Fatalf("Result %+v, %v", res, err) + } + if got := readN(t, int(ptm.Fd()), len(out)); string(got) != out { + t.Fatalf("the terminal got %q, want %q", got, out) + } +} + +// readN reads n bytes from the terminal's master side. It fails the test +// when they do not arrive within 10s. +func readN(t *testing.T, fd, n int) []byte { + t.Helper() + got := make([]byte, 0, n) + deadline := time.Now().Add(10 * time.Second) + for len(got) < n { + wait := int(time.Until(deadline).Milliseconds()) + if wait <= 0 { + t.Fatalf("the terminal got %q of %d bytes", got, n) + } + ready, err := unix.Poll([]unix.PollFd{{Fd: int32(fd), Events: unix.POLLIN}}, wait) + switch { + case err == unix.EINTR || ready == 0: + continue + case err != nil: + t.Fatal(err) + } + m, err := unix.Read(fd, got[len(got):n]) + if err != nil || m == 0 { + t.Fatalf("read the terminal: %v; got %q", err, got) + } + got = got[:len(got)+m] + } + return got +} diff --git a/apps/daemon/internal/processshim/shim_linux.go b/apps/daemon/internal/processshim/shim_linux.go new file mode 100644 index 00000000..d0573c90 --- /dev/null +++ b/apps/daemon/internal/processshim/shim_linux.go @@ -0,0 +1,199 @@ +//go:build linux + +package processshim + +import ( + "bytes" + "fmt" + "net" + "os" + "os/signal" + "runtime" + "syscall" + "unsafe" + + "golang.org/x/sys/unix" +) + +// Forwarded are the signals the shim catches and reports to the relay; the +// broker forwards those the process service declares. They are every signal +// the Go runtime lets a program catch, except CHLD and PIPE, which the shim +// ignores, and URG and PROF, which the runtime uses. +var Forwarded = append([]syscall.Signal{ + unix.SIGHUP, unix.SIGINT, unix.SIGQUIT, unix.SIGABRT, unix.SIGUSR1, unix.SIGUSR2, + unix.SIGALRM, unix.SIGTERM, unix.SIGCONT, unix.SIGTSTP, unix.SIGTTIN, unix.SIGTTOU, + unix.SIGXCPU, unix.SIGXFSZ, unix.SIGVTALRM, unix.SIGWINCH, unix.SIGIO, unix.SIGPWR, +}, realTime(35, 64)...) + +func realTime(first, last syscall.Signal) []syscall.Signal { + var sigs []syscall.Signal + for s := first; s <= last; s++ { + sigs = append(sigs, s) + } + return sigs +} + +// Run runs the shim against the relay at socketPath and returns its exit +// code, or does not return when it re-raises the remote signal. +func Run(socketPath string) int { + caught := make(chan os.Signal, 64) + for _, s := range Forwarded { + // HUP or INT ignored at exec stays ignored, as in a native child. + // The Go runtime replaces other inherited ignores. + if !signal.Ignored(s) { + signal.Notify(caught, s) + } + } + // The Go runtime already ignores SIGURG for the program and uses it for + // preemption, so only CHLD and PIPE are set to ignore here. + signal.Ignore(unix.SIGCHLD, unix.SIGPIPE) + + req, err := request() + if err != nil { + return fail(ExitCannotRun, err.Error()) + } + sock, err := net.DialUnix("unix", nil, &net.UnixAddr{Name: socketPath, Net: "unix"}) + if err != nil { + return fail(ExitCannotRun, "process relay unavailable: "+err.Error()) + } + c := NewConn(sock) + if err := c.SendRequest(req, [3]int{0, 1, 2}); err != nil { + return fail(ExitCannotRun, "process relay: "+err.Error()) + } + m, err := c.ReadMessage() + if err != nil { + return fail(ExitCannotRun, "process relay: "+err.Error()) + } + switch m := m.(type) { + case Result: + if len(m.Message) > 0 { + fmt.Fprintf(os.Stderr, "oac-process-shim: %s\n", m.Message) + } + return exit(m) + case Ack: + default: + return fail(ExitCannotRun, "process relay: unexpected message") + } + + // The relay's copies are now the only references this invocation holds, + // so a reader of the shim's stdout sees EOF when the remote output ends. + os.Stdin.Close() + os.Stdout.Close() + os.Stderr.Close() + + results := make(chan Result, 1) + go func() { + m, err := c.ReadMessage() + if r, ok := m.(Result); ok && err == nil { + results <- r + } + close(results) + }() + for { + select { + case s := <-caught: + // A failed send means the relay is gone; the read reports it. + _ = c.Send(Signal{Number: uint16(s.(syscall.Signal))}) + case r, ok := <-results: + if !ok { + return ExitLost + } + return exit(r) + } + } +} + +func request() (Request, error) { + cwd, err := unix.Getwd() + if err != nil { + return Request{}, fmt.Errorf("working directory: %w", err) + } + umask := unix.Umask(0) + unix.Umask(umask) + r := Request{ + Version: Version, + ExecPath: execPath(), + Cwd: []byte(cwd), + Umask: uint32(umask), + } + for _, a := range os.Args { + r.Argv = append(r.Argv, []byte(a)) + } + for _, v := range os.Environ() { + r.Env = append(r.Env, []byte(v)) + } + if n := len(Frame(r).Payload); n > MaxRequestBytes { + return Request{}, fmt.Errorf("argument list and environment of %d bytes exceed %d", n, MaxRequestBytes) + } + return r, nil +} + +// execPath returns the path the kernel executed, which keeps the directory a +// PATH search chose; argv[0] may be just the name. +func execPath() []byte { + if p := atExecFn(); p != nil { + return p + } + if len(os.Args) > 0 { + return []byte(os.Args[0]) + } + return nil +} + +// atExecFnTag is AT_EXECFN from . +const atExecFnTag = 31 + +func atExecFn() []byte { + auxv, err := unix.Auxv() + if err != nil { + return nil + } + for _, kv := range auxv { + if kv[0] != atExecFnTag { + continue + } + mem, err := os.Open("/proc/self/mem") + if err != nil { + return nil + } + defer mem.Close() + b := make([]byte, unix.PathMax) + n, _ := mem.ReadAt(b, int64(kv[1])) + if i := bytes.IndexByte(b[:n], 0); i > 0 { + return b[:i] + } + return nil + } + return nil +} + +func fail(code int, reason string) int { + fmt.Fprintf(os.Stderr, "oac-process-shim: %s\n", reason) + return code +} + +// exit returns r's code, or ends the process with r's signal the way the +// remote program ended. +func exit(r Result) int { + if r.Signal == 0 { + return int(r.Code) + } + sig := syscall.Signal(r.Signal) + runtime.LockOSThread() + // No core file: the core would be the shim's, not the program's. + _ = unix.Setrlimit(unix.RLIMIT_CORE, &unix.Rlimit{}) + // Go's own handler would dump goroutines for SIGQUIT and similar, so set + // the kernel default directly. + var act [32]byte // struct sigaction with SIG_DFL, no flags, empty mask + _, _, errno := unix.RawSyscall6(unix.SYS_RT_SIGACTION, uintptr(sig), uintptr(unsafe.Pointer(&act)), 0, 8, 0, 0) + if errno == 0 || sig == unix.SIGKILL { + var set unix.Sigset_t + bits := uint(unsafe.Sizeof(set.Val[0])) * 8 + set.Val[(uint(sig)-1)/bits] |= 1 << ((uint(sig) - 1) % bits) + _ = unix.PthreadSigmask(unix.SIG_UNBLOCK, &set, nil) + _ = unix.Tgkill(unix.Getpid(), unix.Gettid(), sig) + } + // The signal does not end the process by default, or it could not be + // raised; report it the way a shell does. + return 128 + int(sig) +} diff --git a/apps/daemon/internal/processshim/shim_other.go b/apps/daemon/internal/processshim/shim_other.go new file mode 100644 index 00000000..ef13c7fc --- /dev/null +++ b/apps/daemon/internal/processshim/shim_other.go @@ -0,0 +1,27 @@ +//go:build !linux + +package processshim + +import ( + "errors" + "fmt" + "os" +) + +// ErrUnsupported is returned on platforms without the process shim. +var ErrUnsupported = errors.New("processshim: unsupported on this platform") + +// Run reports ErrUnsupported and returns ExitCannotRun. +func Run(string) int { + fmt.Fprintf(os.Stderr, "oac-process-shim: %v\n", ErrUnsupported) + return ExitCannotRun +} + +// Relaying reports false. +func Relaying() bool { return false } + +// Relay reports ErrUnsupported and returns 1. +func Relay() int { + fmt.Fprintf(os.Stderr, "oac-process-shim: %v\n", ErrUnsupported) + return 1 +} diff --git a/apps/daemon/internal/processshim/terminal_linux.go b/apps/daemon/internal/processshim/terminal_linux.go new file mode 100644 index 00000000..b2421c82 --- /dev/null +++ b/apps/daemon/internal/processshim/terminal_linux.go @@ -0,0 +1,167 @@ +//go:build linux + +package processshim + +import ( + "sync" + + "golang.org/x/sys/unix" +) + +// terminals coordinates the invocations that share a terminal: the first +// one saves the terminal's mode, each one runs it raw from that mode, and +// the last one to leave restores it. Saving and restoring both happen under +// mu, so an invocation that arrives while the last one leaves saves the +// restored mode, never the raw one. +type terminals struct { + mu sync.Mutex + m map[termKey]*sharedTerminal +} + +// termKey identifies a terminal device. +type termKey struct{ dev, rdev uint64 } + +type sharedTerminal struct { + saved unix.Termios // fixed once created + users int + raw bool +} + +// terminal is the terminal on descriptor 0 of a PTY invocation. It is raw +// while the program runs and restored on every path. It keeps its own +// descriptor, so it outlives the stdin pump's; holding a terminal open has +// no end-of-file effect. +type terminal struct { + fd int + reg *terminals + key termKey + shared *sharedTerminal + + mu sync.Mutex + left bool // the invocation no longer uses the terminal + closed bool +} + +// openTerminal returns the terminal when in and out are both terminals, and +// counts the invocation as one of its users. +func openTerminal(reg *terminals, in, out int) (*terminal, error) { + if !isTerminal(in) || !isTerminal(out) { + return nil, nil + } + var st unix.Stat_t + if err := unix.Fstat(in, &st); err != nil { + return nil, err + } + fd, err := unix.FcntlInt(uintptr(in), unix.F_DUPFD_CLOEXEC, 3) + if err != nil { + return nil, err + } + key := termKey{dev: st.Dev, rdev: st.Rdev} + reg.mu.Lock() + defer reg.mu.Unlock() + s := reg.m[key] + if s == nil { + cur, err := unix.IoctlGetTermios(fd, unix.TCGETS) + if err != nil { + unix.Close(fd) + return nil, err + } + s = &sharedTerminal{saved: *cur} + if reg.m == nil { + reg.m = map[termKey]*sharedTerminal{} + } + reg.m[key] = s + } + s.users++ + return &terminal{fd: fd, reg: reg, key: key, shared: s}, nil +} + +func isTerminal(fd int) bool { + _, err := unix.IoctlGetTermios(fd, unix.TCGETS) + return err == nil +} + +// mode is the terminal's saved mode and current size. +func (t *terminal) mode() *Terminal { + s := t.shared.saved + return &Terminal{ + Size: t.size(), + Iflag: s.Iflag, Oflag: s.Oflag, Cflag: s.Cflag, Lflag: s.Lflag, + Cc: append([]byte(nil), s.Cc[:]...), + } +} + +func (t *terminal) size() WindowSize { + t.mu.Lock() + defer t.mu.Unlock() + if t.closed { + return WindowSize{} + } + ws, err := unix.IoctlGetWinsize(t.fd, unix.TIOCGWINSZ) + if err != nil { + return WindowSize{} + } + return WindowSize{Rows: ws.Row, Cols: ws.Col, XPixels: ws.Xpixel, YPixels: ws.Ypixel} +} + +// makeRaw passes every byte through to the remote terminal, which does the +// line discipline. +func (t *terminal) makeRaw() error { + t.mu.Lock() + defer t.mu.Unlock() + if t.left { + return nil + } + raw := t.shared.saved + raw.Iflag &^= unix.IGNBRK | unix.BRKINT | unix.PARMRK | unix.ISTRIP | unix.INLCR | unix.IGNCR | unix.ICRNL | unix.IXON + raw.Oflag &^= unix.OPOST + raw.Lflag &^= unix.ECHO | unix.ECHONL | unix.ICANON | unix.ISIG | unix.IEXTEN + raw.Cflag &^= unix.CSIZE | unix.PARENB + raw.Cflag |= unix.CS8 + raw.Cc[unix.VMIN], raw.Cc[unix.VTIME] = 1, 0 + t.reg.mu.Lock() + defer t.reg.mu.Unlock() + if err := unix.IoctlSetTermios(t.fd, unix.TCSETS, &raw); err != nil { + return err + } + t.shared.raw = true + return nil +} + +// restore ends the invocation's use of the terminal. The last user restores +// the saved mode. +func (t *terminal) restore() { + if t == nil { + return + } + t.mu.Lock() + defer t.mu.Unlock() + if t.left { + return + } + t.left = true + t.reg.mu.Lock() + defer t.reg.mu.Unlock() + s := t.shared + if s.users--; s.users > 0 { + return + } + delete(t.reg.m, t.key) + if s.raw { + unix.IoctlSetTermios(t.fd, unix.TCSETS, &s.saved) + } +} + +// close restores the terminal and closes its descriptor. +func (t *terminal) close() { + if t == nil { + return + } + t.restore() + t.mu.Lock() + defer t.mu.Unlock() + if !t.closed { + t.closed = true + unix.Close(t.fd) + } +} diff --git a/apps/daemon/internal/processshim/testdata/accept.hex b/apps/daemon/internal/processshim/testdata/accept.hex new file mode 100644 index 00000000..22ab797a --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/accept.hex @@ -0,0 +1,6 @@ +# Accept. +00000008 # PayloadLength 8 +0021 # MessageType: Accept +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 diff --git a/apps/daemon/internal/processshim/testdata/ack.hex b/apps/daemon/internal/processshim/testdata/ack.hex new file mode 100644 index 00000000..058796b5 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/ack.hex @@ -0,0 +1,5 @@ +# An Ack. +00000000 # PayloadLength 0 +0002 # MessageType: Ack +0000 # Flags +0000000000000000 # RequestID 0 diff --git a/apps/daemon/internal/processshim/testdata/close.hex b/apps/daemon/internal/processshim/testdata/close.hex new file mode 100644 index 00000000..2d5f01cf --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/close.hex @@ -0,0 +1,8 @@ +# Close for stdout. +00000011 # PayloadLength 17 +0026 # MessageType: Close +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +01 # FD 1 +0000000000000009 # Seq 9 diff --git a/apps/daemon/internal/processshim/testdata/end.hex b/apps/daemon/internal/processshim/testdata/end.hex new file mode 100644 index 00000000..483ab5ac --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/end.hex @@ -0,0 +1,6 @@ +# End. +00000008 # PayloadLength 8 +0029 # MessageType: End +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 diff --git a/apps/daemon/internal/processshim/testdata/exit.hex b/apps/daemon/internal/processshim/testdata/exit.hex new file mode 100644 index 00000000..f8523eac --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/exit.hex @@ -0,0 +1,12 @@ +# Exit with code 3 once both outputs reach their marks. +00000025 # PayloadLength 37 +0027 # MessageType: Exit +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +0000 # Result Signal 0 +03 # Result Code 3 +00000000 # Result Message "" +00000002 # Marks count 2 +01 0000000000000009 # FD 1 through Seq 9 +02 0000000000000006 # FD 2 through Seq 6 diff --git a/apps/daemon/internal/processshim/testdata/gone.hex b/apps/daemon/internal/processshim/testdata/gone.hex new file mode 100644 index 00000000..9ae82526 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/gone.hex @@ -0,0 +1,6 @@ +# Gone: the shim's connection ended. +00000008 # PayloadLength 8 +0017 # MessageType: Gone +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000002 # ID 2 diff --git a/apps/daemon/internal/processshim/testdata/input.hex b/apps/daemon/internal/processshim/testdata/input.hex new file mode 100644 index 00000000..5ef4715a --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/input.hex @@ -0,0 +1,7 @@ +# Input read from stdin. +0000000f # PayloadLength 15 +0012 # MessageType: Input +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +00000003 68690a # Data "hi\n" diff --git a/apps/daemon/internal/processshim/testdata/input_end.hex b/apps/daemon/internal/processshim/testdata/input_end.hex new file mode 100644 index 00000000..8f82df6e --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/input_end.hex @@ -0,0 +1,6 @@ +# InputEnd at end of file. +00000008 # PayloadLength 8 +0013 # MessageType: InputEnd +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 diff --git a/apps/daemon/internal/processshim/testdata/notice.hex b/apps/daemon/internal/processshim/testdata/notice.hex new file mode 100644 index 00000000..c8bce2c2 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/notice.hex @@ -0,0 +1,7 @@ +# Notice for fd 2. +00000015 # PayloadLength 21 +0028 # MessageType: Notice +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +00000009 6c696e6b206c6f7374 # Message "link lost" diff --git a/apps/daemon/internal/processshim/testdata/open.hex b/apps/daemon/internal/processshim/testdata/open.hex new file mode 100644 index 00000000..af28b0a3 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/open.hex @@ -0,0 +1,20 @@ +# An Open with a terminal. Hex bytes; text after # is a comment. +00000051 # PayloadLength 81 +0011 # MessageType: Open +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +0001 # Request Version 1 +0000000c 2f2e6f61632f62696e2f7368 # ExecPath "/.oac/bin/sh" +00000001 # Argv count 1 +00000002 7368 # "sh" +00000000 # Env count 0 +00000002 2f77 # Cwd "/w" +00000012 # Umask 0o022 +01 # Terminal present +0018 0050 0000 0000 # Size 24 rows, 80 columns +00000500 # Iflag +00000005 # Oflag +000000bf # Cflag +00008a3b # Lflag +00000002 031c # Cc diff --git a/apps/daemon/internal/processshim/testdata/output.hex b/apps/daemon/internal/processshim/testdata/output.hex new file mode 100644 index 00000000..d0008be5 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/output.hex @@ -0,0 +1,9 @@ +# Output for stderr. +00000019 # PayloadLength 25 +0025 # MessageType: Output +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +02 # FD 2 +0000000000000005 # Seq 5 +00000004 6572720a # Data "err\n" diff --git a/apps/daemon/internal/processshim/testdata/read.hex b/apps/daemon/internal/processshim/testdata/read.hex new file mode 100644 index 00000000..141e1b9e --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/read.hex @@ -0,0 +1,7 @@ +# Read granting one Input of up to 64 KiB. +0000000c # PayloadLength 12 +0023 # MessageType: Read +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +00010000 # Max 65536 diff --git a/apps/daemon/internal/processshim/testdata/request.hex b/apps/daemon/internal/processshim/testdata/request.hex new file mode 100644 index 00000000..3e6e4038 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/request.hex @@ -0,0 +1,14 @@ +# A Request. Hex bytes; text after # is a comment. +00000049 # PayloadLength 73 +0001 # MessageType: Request +0000 # Flags +0000000000000000 # RequestID 0 +0001 # Version 1 +0000000d 2f2e6f61632f62696e2f676974 # ExecPath "/.oac/bin/git" +00000002 # Argv count 2 +00000003 676974 # "git" +00000006 737461747573 # "status" +00000001 # Env count 1 +0000000c 484f4d453d2f686f6d652f75 # "HOME=/home/u" +00000005 2f776f726b # Cwd "/work" +00000012 # Umask 0o022 diff --git a/apps/daemon/internal/processshim/testdata/result_code.hex b/apps/daemon/internal/processshim/testdata/result_code.hex new file mode 100644 index 00000000..51455473 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/result_code.hex @@ -0,0 +1,8 @@ +# A Result with exit code 3. +00000007 # PayloadLength 7 +0004 # MessageType: Result +0000 # Flags +0000000000000000 # RequestID 0 +0000 # Signal 0 +03 # Code 3 +00000000 # Message "" diff --git a/apps/daemon/internal/processshim/testdata/result_refused.hex b/apps/daemon/internal/processshim/testdata/result_refused.hex new file mode 100644 index 00000000..1ecc11a9 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/result_refused.hex @@ -0,0 +1,8 @@ +# A Result sent instead of the Ack. +00000010 # PayloadLength 16 +0004 # MessageType: Result +0000 # Flags +0000000000000000 # RequestID 0 +0000 # Signal 0 +7f # Code 127 +00000009 6e6f7420666f756e64 # Message "not found" diff --git a/apps/daemon/internal/processshim/testdata/result_signal.hex b/apps/daemon/internal/processshim/testdata/result_signal.hex new file mode 100644 index 00000000..c09a4d2d --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/result_signal.hex @@ -0,0 +1,8 @@ +# A Result for a program ended by SIGTERM. +00000007 # PayloadLength 7 +0004 # MessageType: Result +0000 # Flags +0000000000000000 # RequestID 0 +000f # Signal 15 +00 # Code 0 +00000000 # Message "" diff --git a/apps/daemon/internal/processshim/testdata/signal.hex b/apps/daemon/internal/processshim/testdata/signal.hex new file mode 100644 index 00000000..eae4900c --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/signal.hex @@ -0,0 +1,6 @@ +# A Signal reporting SIGINT. +00000002 # PayloadLength 2 +0003 # MessageType: Signal +0000 # Flags +0000000000000000 # RequestID 0 +0002 # Number 2 diff --git a/apps/daemon/internal/processshim/testdata/signaled.hex b/apps/daemon/internal/processshim/testdata/signaled.hex new file mode 100644 index 00000000..71cd3fb8 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/signaled.hex @@ -0,0 +1,9 @@ +# Signaled for SIGWINCH on a terminal. +00000013 # PayloadLength 19 +0016 # MessageType: Signaled +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +001c # Number 28 +01 # Size present +001e 0064 0000 0000 # Size 30 rows, 100 columns diff --git a/apps/daemon/internal/processshim/testdata/started.hex b/apps/daemon/internal/processshim/testdata/started.hex new file mode 100644 index 00000000..bb2179ca --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/started.hex @@ -0,0 +1,6 @@ +# Started. +00000008 # PayloadLength 8 +0022 # MessageType: Started +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 diff --git a/apps/daemon/internal/processshim/testdata/stop_input.hex b/apps/daemon/internal/processshim/testdata/stop_input.hex new file mode 100644 index 00000000..fbbb84f6 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/stop_input.hex @@ -0,0 +1,6 @@ +# StopInput. +00000008 # PayloadLength 8 +0024 # MessageType: StopInput +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 diff --git a/apps/daemon/internal/processshim/testdata/write_failed.hex b/apps/daemon/internal/processshim/testdata/write_failed.hex new file mode 100644 index 00000000..80eb95b5 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/write_failed.hex @@ -0,0 +1,9 @@ +# WriteFailed with EACCES. +00000015 # PayloadLength 21 +0015 # MessageType: WriteFailed +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +01 # FD 1 +0000000000000008 # Seq 8 +0000000d # Errno 13 diff --git a/apps/daemon/internal/processshim/testdata/written.hex b/apps/daemon/internal/processshim/testdata/written.hex new file mode 100644 index 00000000..1371cb70 --- /dev/null +++ b/apps/daemon/internal/processshim/testdata/written.hex @@ -0,0 +1,8 @@ +# Written for an Output on fd 1. +00000011 # PayloadLength 17 +0014 # MessageType: Written +0000 # Flags +0000000000000000 # RequestID 0 +0000000000000001 # ID 1 +01 # FD 1 +0000000000000007 # Seq 7 diff --git a/apps/daemon/internal/sessionview/build_linux.go b/apps/daemon/internal/sessionview/build_linux.go index ae5c4131..dbab1cb3 100644 --- a/apps/daemon/internal/sessionview/build_linux.go +++ b/apps/daemon/internal/sessionview/build_linux.go @@ -5,18 +5,23 @@ package sessionview import ( "errors" "fmt" + "slices" "strings" "golang.org/x/sys/unix" "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/agent" + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" ) // builder mounts the local pieces onto the world, at the paths the world presents the mountpoints at. Every target is resolved beneath its parent mount without following symlinks, so the sandbox cannot redirect a mount. type builder struct { root int // the world root, an O_PATH fd targets map[string]string // each mountpoint's view path to where the world presents it - fds []int + fds []int // what only the build needs + + // What the launcher keeps: the view's /proc, and the relay's listening socket or -1. + proc, listener int } // devNodes are bound from the host's /dev into the view's /dev. @@ -55,15 +60,21 @@ func (b *builder) build(spec *launchSpec) error { return err } } - shim := -1 - if len(spec.Shim.Names) > 0 || len(spec.Shim.Paths) > 0 { + shim, names := -1, spec.Shim.Names + if spec.Shim.declared() { if shim, err = b.source(spec.Shim.Binary); err != nil { return err } + names = append(slices.Clone(names), processshim.RelayName) } - if err := b.shimDir(spec.Shim.Names, shim); err != nil { + if err := b.shimDir(names, shim); err != nil { return err } + if spec.Shim.declared() { + if err := b.runDir(spec.UID, spec.GID); err != nil { + return err + } + } for _, o := range spec.Overlays { at, err := b.at(o.Path) if err != nil { @@ -90,12 +101,10 @@ func (b *builder) build(spec *launchSpec) error { if err != nil { return err } - proc, err := newFS("proc", nil, attrNoSuid|attrNoDev|attrNoExec) - if err != nil { + if b.proc, err = newFS("proc", nil, attrNoSuid|attrNoDev|attrNoExec); err != nil { return err } - defer unix.Close(proc) - if err := b.attach(proc, b.root, at, true); err != nil { + if err := b.attach(b.proc, b.root, at, true); err != nil { return err } return b.dev() @@ -134,6 +143,43 @@ func (b *builder) shimDir(names []string, shim int) error { return readOnly(mnt, dir) } +// runDir presents the relay's listening socket at processshim.SocketPath on a read-only tmpfs, so that the process can connect to it but not replace it. Only the process's user may connect. The socket is created through the mount, never through a path the world resolves. +func (b *builder) runDir(uid, gid uint32) error { + dir, err := b.at(agent.ViewPrivateRoot + "/" + agent.ViewRunName) + if err != nil { + return err + } + mnt, err := newFS("tmpfs", [][2]string{{"mode", "0755"}, {"size", "64k"}}, attrNoSuid|attrNoDev|attrNoExec) + if err != nil { + return err + } + defer unix.Close(mnt) + if err := b.attach(mnt, b.root, dir, true); err != nil { + return err + } + fail := func(op string, err error) error { + return &Error{Kind: ErrLauncher, Op: op, Path: processshim.SocketPath, Err: err} + } + ln, err := unix.Socket(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + if err != nil { + return fail("socket", err) + } + b.listener = ln + if err := unix.Bind(ln, &unix.SockaddrUnix{Name: fmt.Sprintf("/proc/self/fd/%d/%s", mnt, processshim.SocketName)}); err != nil { + return fail("bind", err) + } + if err := unix.Fchownat(mnt, processshim.SocketName, int(uid), int(gid), unix.AT_SYMLINK_NOFOLLOW); err != nil { + return fail("chown", err) + } + if err := unix.Fchmodat(mnt, processshim.SocketName, 0o600, 0); err != nil { + return fail("chmod", err) + } + if err := unix.Listen(ln, unix.SOMAXCONN); err != nil { + return fail("listen", err) + } + return readOnly(mnt, dir) +} + // dev builds a minimal read-only /dev with host device nodes, a new devpts instance and a noexec /dev/shm. func (b *builder) dev() error { at, err := b.at(agent.ViewDevRoot) @@ -204,13 +250,30 @@ func (b *builder) keep(fd int) int { return fd } -func (b *builder) close() { +// closeBuild closes what only the build needs. +func (b *builder) closeBuild() { for _, fd := range b.fds { unix.Close(fd) } b.fds, b.root = nil, -1 } +func (b *builder) closeListener() { + if b.listener >= 0 { + unix.Close(b.listener) + b.listener = -1 + } +} + +func (b *builder) close() { + b.closeBuild() + b.closeListener() + if b.proc >= 0 { + unix.Close(b.proc) + b.proc = -1 + } +} + // bind mounts a clone of src at the absolute view path under the world root. func (b *builder) bind(src, root int, view string, attr *unix.MountAttr) error { return b.bindAt(src, root, strings.TrimPrefix(view, "/"), view, attr) diff --git a/apps/daemon/internal/sessionview/control_linux.go b/apps/daemon/internal/sessionview/control_linux.go index 65fe9065..37c6e2c2 100644 --- a/apps/daemon/internal/sessionview/control_linux.go +++ b/apps/daemon/internal/sessionview/control_linux.go @@ -17,13 +17,14 @@ import ( "golang.org/x/sys/unix" ) -// Launcher file descriptors. The spec travels over a pipe so that nothing about the view appears in argv or the environment. +// Launcher file descriptors. The spec travels over a pipe so that nothing about the view appears in argv or the environment. relayFD, the relay's end of its broker connection, is open only when the spec declares a shim. const ( specFD = 3 controlFD = 4 stdinFD = 5 stdoutFD = 6 stderrFD = 7 + relayFD = 8 ) // launchSpec is what the daemon sends the launcher: the Spec without its callbacks and files. diff --git a/apps/daemon/internal/sessionview/doc.go b/apps/daemon/internal/sessionview/doc.go index 1e296a51..4966a9c7 100644 --- a/apps/daemon/internal/sessionview/doc.go +++ b/apps/daemon/internal/sessionview/doc.go @@ -6,5 +6,9 @@ // // The daemon calls [Init] first thing in main. [Start] re-executes the daemon binary as the launcher, which becomes PID 1 of the view: it builds the view, starts the process, delivers signals to every process in the view, reaps orphans and exits with the process status. Once the process has exited, Signal reports [ErrExited] and delivers nothing. When the process exits while others remain, the launcher sends them TERM unless one already went to the view, and waits for them up to [Process].Grace from the first TERM. Its exit kills what remains and tears the view down. // +// A view that declares a [Shim] also runs the Session's process relay, the shim binary in relay mode (package processshim). The launcher starts it before the process, under the same restrictions and as the same user, with a listening socket at processshim.SocketPath on a read-only mount and its end of a socket pair whose other end is [View.Relay]. The relay is the only process that receives the descriptors a shim hands over; the broker outside the view holds only its end of the pair. The launcher's drain ignores the relay, which ends with the view. +// +// Teardown never waits unconditionally. Once the launcher has exited, the view shuts its end of the relay connection down, stops the world server, which ends the requests still pending on the view's FUSE connection so that a process blocked on the world can exit, and waits for the world server and the view's processes together for up to 30 seconds. Past that, Wait and Close report [ErrCleanup]. +// // The package works only on Linux. Elsewhere [Start] returns [ErrUnsupported]. package sessionview diff --git a/apps/daemon/internal/sessionview/errors.go b/apps/daemon/internal/sessionview/errors.go index e7dcc5ba..e2bb26ee 100644 --- a/apps/daemon/internal/sessionview/errors.go +++ b/apps/daemon/internal/sessionview/errors.go @@ -21,6 +21,8 @@ var ( ErrExec = errors.New("sessionview: process start failed") ErrLauncher = errors.New("sessionview: launcher failed") ErrClosed = errors.New("sessionview: view closed") + // ErrCleanup reports a teardown that did not finish within its bound: the view's processes or the world server were still running. What remains finishes in the background if it can. + ErrCleanup = errors.New("sessionview: cleanup incomplete") // ErrExited is Signal's result once the process has exited. It also matches os.ErrProcessDone. ErrExited = errors.New("sessionview: process exited") ) diff --git a/apps/daemon/internal/sessionview/launcher_linux.go b/apps/daemon/internal/sessionview/launcher_linux.go index 6a8bd752..f1ed252e 100644 --- a/apps/daemon/internal/sessionview/launcher_linux.go +++ b/apps/daemon/internal/sessionview/launcher_linux.go @@ -8,12 +8,15 @@ import ( "os" "os/signal" "runtime" + "strconv" "strings" "sync" "syscall" "time" "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" ) // Init runs the launcher when the process was started as one, and then never returns. The daemon calls it first thing in main. @@ -33,6 +36,9 @@ type launcher struct { proceedOnce sync.Once targets map[string]string // set before proceed closes + proc int // the view's /proc + relay int // the relay's pid, or 0 + // mu orders signals against the process's exit. running holds from the process's start until it is reaped; termAt is when TERM first went to the view. mu sync.Mutex running bool @@ -70,18 +76,28 @@ func (l *launcher) run() (int, error) { return 0, err } <-l.proceed - b := &builder{root: -1, targets: l.targets} + b := &builder{root: -1, proc: -1, listener: -1, targets: l.targets} defer b.close() if err := b.build(spec); err != nil { return 0, err } + l.proc = b.proc if err := switchRoot(b.root); err != nil { return 0, err } if err := restrict(); err != nil { return 0, err } - b.close() + b.closeBuild() + if b.listener >= 0 { + // The launcher keeps neither the listener nor the broker connection, so the relay's end is theirs alone. + l.relay, err = startRelay(spec, b.listener) + b.closeListener() + unix.Close(relayFD) + if err != nil { + return 0, err + } + } pid, err := startProcess(spec) if err != nil { return 0, err @@ -157,7 +173,7 @@ func (l *launcher) forwardSignals() { }() } -// drain gives the processes left after the process exits TERM and up to grace from that TERM to exit, reaping them. When TERM already went to the view, they keep what remains of the grace and get no second TERM. The launcher's exit then kills whatever remains. +// drain gives the processes left after the process exits TERM and up to grace from that TERM to exit, reaping them, and ends once only the relay remains. When TERM already went to the view, they keep what remains of the grace and get no second TERM. The launcher's exit then kills whatever remains, the relay too. func (l *launcher) drain(grace time.Duration) { l.mu.Lock() termAt := l.termAt @@ -171,7 +187,8 @@ func (l *launcher) drain(grace time.Duration) { reaped := make(chan struct{}) go func() { defer close(reaped) - for { + // The last process other than the relay to exit is a child of the launcher by then, so its exit ends the Wait4. + for l.othersRemain() { if _, err := unix.Wait4(-1, nil, 0, nil); err != nil && err != unix.EINTR { return } @@ -185,6 +202,26 @@ func (l *launcher) drain(grace time.Duration) { } } +// othersRemain reports whether a process other than the launcher and the relay remains in the view. It errs on yes. +func (l *launcher) othersRemain() bool { + fd, err := unix.Openat(l.proc, ".", unix.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC, 0) + if err != nil { + return true + } + dir := os.NewFile(uintptr(fd), "proc") + defer dir.Close() + names, err := dir.Readdirnames(-1) + if err != nil { + return true + } + for _, n := range names { + if pid, err := strconv.Atoi(n); err == nil && pid != 1 && pid != l.relay { + return true + } + } + return false +} + // reap collects every child, since orphans in the view reparent to PID 1, until the process exits. func (l *launcher) reap(pid int) (int, error) { for { @@ -253,6 +290,28 @@ func loopbackUp() error { return unix.IoctlIfreq(fd, unix.SIOCSIFFLAGS, ifr) } +// startRelay starts the Session's process relay as the process's user, with the broker connection and the listening socket at the descriptors processshim gives them, no standard descriptors and an empty environment. +func startRelay(spec *launchSpec, listener int) (int, error) { + files := make([]uintptr, max(processshim.RelayBrokerFD, processshim.RelayListenerFD)+1) + for i := range files { + files[i] = ^uintptr(0) // closed + } + files[processshim.RelayBrokerFD], files[processshim.RelayListenerFD] = relayFD, uintptr(listener) + pid, err := syscall.ForkExec(processshim.RelayPath, processshim.RelayArgs, &syscall.ProcAttr{ + Dir: "/", + Env: []string{}, + Files: files, + Sys: &syscall.SysProcAttr{ + Setsid: true, + Credential: &syscall.Credential{Uid: spec.UID, Gid: spec.GID, Groups: spec.Groups}, + }, + }) + if err != nil { + return 0, &Error{Kind: ErrExec, Op: "exec relay", Path: processshim.RelayPath, Err: err} + } + return pid, nil +} + func startProcess(spec *launchSpec) (int, error) { pid, err := syscall.ForkExec(spec.Path, spec.Args, &syscall.ProcAttr{ Dir: spec.Dir, diff --git a/apps/daemon/internal/sessionview/spec.go b/apps/daemon/internal/sessionview/spec.go index 17204a52..6d3f84d8 100644 --- a/apps/daemon/internal/sessionview/spec.go +++ b/apps/daemon/internal/sessionview/spec.go @@ -10,6 +10,7 @@ import ( "time" "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/agent" + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" ) // Spec declares one view and the process it runs. @@ -31,7 +32,7 @@ type World func(ctx context.Context, dev *os.File, mount WorldMount) (WorldServe // WorldServer is a running world. type WorldServer interface { - // Stop returns once serving has ended. sessionview calls it after the view has exited, when the kernel has aborted the connection. + // Stop ends serving and returns within a bound of its own. sessionview calls it once the launcher has exited, while the view's processes may still be ending, and waits for it and for them together. A process blocked on a request the world has not answered cannot exit, and it keeps the view's mount alive, so Stop must end every request still pending rather than only wait for the view to end. Stop() error } @@ -84,7 +85,7 @@ type Overlay struct { Exec bool } -// Shim presents a static binary at agent.ViewPrivateRoot/agent.ViewShimName/ for each name and binds it over each absolute view path. +// Shim presents a static binary at agent.ViewPrivateRoot/agent.ViewShimName/ for each name and binds it over each absolute view path. A Spec that declares a shim also runs the binary as the Session's process relay: it is presented as processshim.RelayName too, and the launcher creates the relay's socket at processshim.SocketPath and starts the relay as the process's user, before the process. [View.Relay] is the broker's end of the relay's connection. type Shim struct { Binary string Names []string @@ -130,7 +131,7 @@ func (s *Spec) validate() error { } names := map[string]bool{} for _, d := range s.Private { - if !isComponent(d.Name) || d.Name == agent.ViewShimName || names[d.Name] { + if !isComponent(d.Name) || d.Name == agent.ViewShimName || d.Name == agent.ViewRunName || names[d.Name] { return invalid("private directory name %q", d.Name) } names[d.Name] = true @@ -153,7 +154,7 @@ func (s *Spec) validate() error { } claimed = append(claimed, o.Path) } - if len(s.Shim.Names) > 0 || len(s.Shim.Paths) > 0 { + if s.Shim.declared() { if info, err := hostSource(s.Shim.Binary); err != nil { return err } else if info.IsDir() { @@ -162,7 +163,7 @@ func (s *Spec) validate() error { } names = map[string]bool{} for _, n := range s.Shim.Names { - if !isComponent(n) || names[n] { + if !isComponent(n) || n == processshim.RelayName || names[n] { return invalid("shim name %q", n) } names[n] = true @@ -183,6 +184,11 @@ func (s *Spec) validate() error { return s.Process.validate() } +// declared reports whether the spec declares a shim, and so a relay. +func (s *Shim) declared() bool { + return len(s.Names) > 0 || len(s.Paths) > 0 +} + func (p *Process) validate() error { if !isViewAbs(p.Path) || !isViewAbs(p.Dir) { return invalid("process path %q and directory %q must be absolute and clean", p.Path, p.Dir) diff --git a/apps/daemon/internal/sessionview/view_linux.go b/apps/daemon/internal/sessionview/view_linux.go index 7303f2a4..1892608c 100644 --- a/apps/daemon/internal/sessionview/view_linux.go +++ b/apps/daemon/internal/sessionview/view_linux.go @@ -10,9 +10,11 @@ import ( "io" "os" "os/exec" + "strings" "sync" "sync/atomic" "syscall" + "time" "golang.org/x/sys/unix" @@ -30,6 +32,12 @@ var ( const fuseMountFlags = unix.MS_NOSUID | unix.MS_NODEV | unix.MS_NOEXEC +// closeWait bounds the teardown once the launcher has exited: the world server's Stop and the end of the view's processes, which Stop may have to unblock. +var closeWait = 30 * time.Second + +// exitWait bounds how long a lost launcher's exit status is awaited for an error message. +const exitWait = time.Second + // View is a running view. type View struct { cmd *exec.Cmd @@ -39,17 +47,19 @@ type View struct { dev *os.File staging string pipes [3]*os.File + relay *os.File // the broker's end of the relay connection - reapOnce sync.Once - reapErr error - signalMu sync.Mutex - signaled chan bool - exited atomic.Bool - closing atomic.Bool - closeOnce sync.Once - done chan struct{} - exit Exit - err error + waited chan struct{} // closed once the launcher has been reaped + waitErr error + cleanupErr error // set before done closes + signalMu sync.Mutex + signaled chan bool + exited atomic.Bool + closing atomic.Bool + closeOnce sync.Once + done chan struct{} + exit Exit + err error } // Start builds a view for spec and starts its process. ctx bounds only the construction. @@ -62,8 +72,7 @@ func Start(ctx context.Context, spec Spec) (*View, error) { } v := &View{done: make(chan struct{}), signaled: make(chan bool, 1)} if err := v.launch(&spec); err != nil { - v.abort() - return nil, err + return nil, v.abort(err) } stop := context.AfterFunc(ctx, func() { _ = v.cmd.Process.Kill() }) err := v.handshake(ctx, &spec) @@ -72,8 +81,7 @@ func Start(ctx context.Context, spec Spec) (*View, error) { err = errors.Join(&Error{Kind: ErrLauncher, Op: "start", Err: ctx.Err()}, err) } if err != nil { - v.abort() - return nil, err + return nil, v.abort(err) } go v.watch() return v, nil @@ -112,12 +120,25 @@ func (v *View) launch(spec *Spec) error { specW.Close() return &Error{Kind: ErrLauncher, Op: "control", Err: err} } + files := []*os.File{specR, ctlChild, child[0], child[1], child[2]} + if spec.Shim.declared() { + // A stream socket, so that the broker can read it without accepting descriptors. + pair, err := unix.Socketpair(unix.AF_UNIX, unix.SOCK_STREAM|unix.SOCK_CLOEXEC, 0) + if err != nil { + specW.Close() + return &Error{Kind: ErrLauncher, Op: "socketpair", Err: err} + } + v.relay = os.NewFile(uintptr(pair[0]), "sessionview-relay") + relayChild := os.NewFile(uintptr(pair[1]), "sessionview-relay") + defer relayChild.Close() + files = append(files, relayChild) + } v.cmd = &exec.Cmd{ Path: "/proc/self/exe", Args: []string{launcherArg0}, Env: []string{}, Stderr: os.Stderr, - ExtraFiles: []*os.File{specR, ctlChild, child[0], child[1], child[2]}, + ExtraFiles: files, // Cloning into the namespaces, rather than unsharing later, puts every runtime thread of the launcher in them and makes it PID 1 of the view. SysProcAttr: &syscall.SysProcAttr{ Cloneflags: syscall.CLONE_NEWNS | syscall.CLONE_NEWNET | syscall.CLONE_NEWPID, @@ -129,6 +150,11 @@ func (v *View) launch(spec *Spec) error { v.cmd = nil return &Error{Kind: ErrLauncher, Op: "start", Err: err} } + v.waited = make(chan struct{}) + go func() { + v.waitErr = v.cmd.Wait() + close(v.waited) + }() ls := &launchSpec{ Staging: v.staging, Private: spec.Private, Overlays: spec.Overlays, Shim: spec.Shim, Path: spec.Process.Path, Args: spec.Process.Args, Env: spec.Process.Env, Dir: spec.Process.Dir, @@ -142,7 +168,7 @@ func (v *View) launch(spec *Spec) error { return nil } -// stdio returns the launcher's ends of the process's stdio, creating pipes where the spec gives no file. +// stdio returns the launcher's ends of the process's stdio, creating pipes where the spec gives no file. The pipes belong to the process's user, so that the relay can open them as its own. func (v *View) stdio(p *Process) ([3]*os.File, error) { child := [3]*os.File{p.Stdin, p.Stdout, p.Stderr} for i := range child { @@ -150,6 +176,12 @@ func (v *View) stdio(p *Process) ([3]*os.File, error) { continue } r, w, err := os.Pipe() + if err == nil { + if err = r.Chown(int(p.UID), int(p.GID)); err != nil { + r.Close() + w.Close() + } + } if err != nil { for j := range i { if v.pipes[j] != nil { @@ -230,37 +262,77 @@ func targetsOf(mps []Mountpoint, p Presentation) (map[string]string, error) { return targets, nil } -// lost reports a launcher that stopped talking, with its exit status when it has exited. +// lost reports a launcher that stopped talking, with its exit status when it exits promptly. func (v *View) lost(op string, err error) error { if errors.Is(err, io.EOF) { - if werr := v.reap(); werr != nil { - err = werr + select { + case <-v.waited: + if v.waitErr != nil { + err = v.waitErr + } + case <-time.After(exitWait): } } return &Error{Kind: ErrLauncher, Op: op, Err: err} } -func (v *View) reap() error { - v.reapOnce.Do(func() { v.reapErr = v.cmd.Wait() }) - return v.reapErr -} - -// abort undoes a failed Start. -func (v *View) abort() { +// abort undoes a failed Start. It returns err, joined with ErrCleanup when the teardown did not finish. +func (v *View) abort(err error) error { if v.cmd != nil { _ = v.cmd.Process.Kill() - _ = v.reap() } - _ = v.teardown() - for _, p := range v.pipes { - if p != nil { - p.Close() + _, _ = v.teardown() + v.closePipes() + if v.cleanupErr != nil { + return errors.Join(err, v.cleanupErr) + } + return err +} + +// teardown releases the view once its launcher is exiting or never started. It disconnects the relay, so that the broker stops using it, and stops the world server, which ends the requests still pending on the view's FUSE connection so that a process blocked on the world can exit. It waits for the launcher and the world server together up to closeWait; past that it sets cleanupErr, and what still runs finishes in the background. It returns the launcher's wait error and the teardown's errors. +func (v *View) teardown() (werr, err error) { + if v.relay != nil { + shutdown(v.relay) + } + waited := v.waited + stopped := make(chan error, 1) + go func() { stopped <- v.stopWorld() }() + timer := time.NewTimer(closeWait) + defer timer.Stop() + var errs []error + for waited != nil || stopped != nil { + select { + case <-waited: + werr, waited = v.waitErr, nil + case serr := <-stopped: + errs, stopped = append(errs, serr), nil + case <-timer.C: + var running []string + if waited != nil { + running = append(running, "the view's processes") + } + if stopped != nil { + running = append(running, "the world server") + } + v.cleanupErr = &Error{Kind: ErrCleanup, Op: "teardown", Err: fmt.Errorf("%s still running after %v", strings.Join(running, " and "), closeWait)} + errs = append(errs, v.cleanupErr) + waited, stopped = nil, nil } } + if v.relay != nil { + v.relay.Close() + } + if v.ctl != nil { + v.ctl.close() + } + if v.staging != "" { + os.Remove(v.staging) + } + return werr, errors.Join(errs...) } -// teardown releases what the view held once the launcher has exited. -func (v *View) teardown() error { +// stopWorld stops the world server, then closes the /dev/fuse connection it served. +func (v *View) stopWorld() error { var err error if v.world != nil { if serr := v.world.Stop(); serr != nil { @@ -270,13 +342,22 @@ func (v *View) teardown() error { if v.dev != nil { v.dev.Close() } - if v.ctl != nil { - v.ctl.close() + return err +} + +// shutdown ends both directions of a socket, whoever else holds it. +func shutdown(f *os.File) { + if c, err := f.SyscallConn(); err == nil { + c.Control(func(fd uintptr) { unix.Shutdown(int(fd), unix.SHUT_RDWR) }) } - if v.staging != "" { - os.Remove(v.staging) +} + +func (v *View) closePipes() { + for _, p := range v.pipes { + if p != nil { + p.Close() + } } - return err } func (v *View) watch() { @@ -302,8 +383,7 @@ func (v *View) watch() { failed = m.Fail.err() } } - werr := v.reap() - terr := v.teardown() + werr, terr := v.teardown() switch { case exit != nil: v.exit = *exit @@ -319,7 +399,7 @@ func (v *View) watch() { } } -// Wait returns how the process ended once the view is torn down and the world server has stopped. It returns ErrClosed when Close ended the view first. +// Wait returns how the process ended once the view is torn down and the world server has stopped, or once the teardown's bound expired, which it reports with ErrCleanup. It returns ErrClosed when Close ended the view first. func (v *View) Wait() (Exit, error) { <-v.done return v.exit, v.err @@ -358,21 +438,20 @@ func (v *View) Signal(sig syscall.Signal) error { } } -// Close kills the view, waits for its teardown and closes the pipes it created. +// Close kills the view, waits for its teardown and closes the pipes it created. It returns ErrCleanup when the teardown did not finish within its bound, and nil otherwise. func (v *View) Close() error { v.closeOnce.Do(func() { v.closing.Store(true) _ = v.cmd.Process.Kill() <-v.done - for _, p := range v.pipes { - if p != nil { - p.Close() - } - } + v.closePipes() }) - return nil + return v.cleanupErr } +// Relay is the broker's end of the relay's connection, or nil when the spec declares no shim. processbroker.Start takes a duplicate of it. The view shuts the connection down as it ends, so a broker still running then sees the relay lost. +func (v *View) Relay() *os.File { return v.relay } + // Stdin is the write end of the process's stdin pipe, or nil when the spec gave a file. func (v *View) Stdin() *os.File { return v.pipes[0] } @@ -387,8 +466,11 @@ func (s *Spec) mountpoints() []Mountpoint { for _, d := range s.Private { m = append(m, Mountpoint{Path: agent.ViewPrivateRoot + "/" + d.Name, Dir: true}) } + m = append(m, Mountpoint{Path: agent.ViewPrivateRoot + "/" + agent.ViewShimName, Dir: true}) + if s.Shim.declared() { + m = append(m, Mountpoint{Path: agent.ViewPrivateRoot + "/" + agent.ViewRunName, Dir: true}) + } m = append(m, - Mountpoint{Path: agent.ViewPrivateRoot + "/" + agent.ViewShimName, Dir: true}, Mountpoint{Path: agent.ViewProcRoot, Dir: true}, Mountpoint{Path: agent.ViewDevRoot, Dir: true}, ) diff --git a/apps/daemon/internal/sessionview/view_linux_test.go b/apps/daemon/internal/sessionview/view_linux_test.go index f68a2b97..870cae15 100644 --- a/apps/daemon/internal/sessionview/view_linux_test.go +++ b/apps/daemon/internal/sessionview/view_linux_test.go @@ -21,6 +21,7 @@ import ( "slices" "strconv" "strings" + "sync" "syscall" "testing" "time" @@ -29,6 +30,8 @@ import ( "github.com/hanwen/go-fuse/v2/fuse" "golang.org/x/net/bpf" "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" ) // The view tests need root with CAP_SYS_ADMIN and CAP_NET_ADMIN, /dev/fuse and no AppArmor confinement. Run them in a throwaway container: @@ -44,9 +47,12 @@ const ( brokerAddr = "127.0.0.1:7070" ) -// The test binary is also the Harness, the shim and the world binary inside the view. +// The test binary is also the Harness, the shim, the relay and the world binary inside the view. func TestMain(m *testing.M) { Init() + if processshim.Relaying() { + os.Exit(processshim.Relay()) + } switch filepath.Base(os.Args[0]) { case "sh", "env": fmt.Println(shimMarker) @@ -168,6 +174,45 @@ func TestViewDescendantsKeepTheGrace(t *testing.T) { } } +// TestTeardownIsBounded checks that a world server that never ends the request the view's process is blocked on fails the teardown with ErrCleanup within the bound instead of hanging it. +func TestTeardownIsBounded(t *testing.T) { + requireView(t) + f := newFixture(t) + writeFile(t, filepath.Join(f.world, "data", "hang"), "") + w := &hangWorld{loopbackWorld: loopbackWorld{dir: f.world}, hung: make(chan struct{}), release: make(chan struct{})} + saved := closeWait + closeWait = time.Second + defer func() { closeWait = saved }() + spec := f.spec(&w.loopbackWorld, "hang") + spec.World = w.serve + v, err := Start(context.Background(), spec) + if err != nil { + t.Fatalf("Start: %v", err) + } + select { + case <-w.hung: + case <-time.After(10 * time.Second): + t.Fatal("the process never wrote to the world") + } + started := time.Now() + if err := v.Close(); !errors.Is(err, ErrCleanup) { + t.Errorf("Close = %v, want ErrCleanup", err) + } + if elapsed := time.Since(started); elapsed > 5*time.Second { + t.Errorf("Close took %v with a teardown bound of %v", elapsed, closeWait) + } + if _, err := v.Wait(); !errors.Is(err, ErrClosed) || !errors.Is(err, ErrCleanup) { + t.Errorf("Wait = %v, want ErrClosed and ErrCleanup", err) + } + // Once the world answers, what was left finishes. + close(w.release) + select { + case <-w.served: + case <-time.After(10 * time.Second): + t.Error("the view outlived the world's answer") + } +} + // TestViewRefusesSymlinkedMountpoint checks that the launcher still refuses a symlink on the way to a target, so a world that reports a target it did not resolve cannot redirect a mount. func TestViewRefusesSymlinkedMountpoint(t *testing.T) { requireView(t) @@ -235,6 +280,8 @@ func TestStartRejectsInvalidSpec(t *testing.T) { "relative overlay": {Overlays: []Overlay{{Path: "etc/resolv.conf", Source: "/etc/hosts"}}}, "writable executable private": {Private: []PrivateDir{{Name: "home", HostDir: t.TempDir(), Writable: true, Exec: true}}}, "missing staging parent": {StagingParent: filepath.Join(t.TempDir(), "missing")}, + "private run directory": {Private: []PrivateDir{{Name: "run", HostDir: t.TempDir()}}}, + "shim named as the relay": {Shim: Shim{Binary: "/bin/true", Names: []string{processshim.RelayName}}}, } { if spec.StagingParent == "" { spec.StagingParent = t.TempDir() @@ -258,7 +305,7 @@ func requireView(t *testing.T) { } type fixture struct { - self, world, harness, home, run, overlay, staging string + self, world, harness, home, overlay, staging string } // newFixture lays out a world with the mountpoints the real world frontend presents synthetically, plus the local sources. @@ -273,7 +320,6 @@ func newFixture(t *testing.T) *fixture { world: filepath.Join(base, "world"), harness: filepath.Join(base, "harness"), home: filepath.Join(base, "home"), - run: filepath.Join(base, "run"), overlay: filepath.Join(base, "overlay"), staging: filepath.Join(base, "staging"), } @@ -286,11 +332,9 @@ func newFixture(t *testing.T) *fixture { copyFile(t, self, filepath.Join(f.world, "bin", "worldbin")) mkdir(t, f.harness) copyFile(t, self, filepath.Join(f.harness, "harness")) - for _, d := range []string{f.home, f.run} { - mkdir(t, d) - if err := os.Chown(d, viewID, viewID); err != nil { - t.Fatal(err) - } + mkdir(t, f.home) + if err := os.Chown(f.home, viewID, viewID); err != nil { + t.Fatal(err) } mkdir(t, f.overlay) mkdir(t, f.staging) @@ -304,7 +348,6 @@ func (f *fixture) spec(w *loopbackWorld, mode string, env ...string) Spec { Private: []PrivateDir{ {Name: "harness", HostDir: f.harness, Exec: true}, {Name: "home", HostDir: f.home, Writable: true}, - {Name: "run", HostDir: f.run, Writable: true}, }, Overlays: []Overlay{{Path: "/etc/oac-overlay", Source: f.overlay}}, Shim: Shim{ @@ -332,13 +375,16 @@ type loopbackWorld struct { } func (w *loopbackWorld) serve(_ context.Context, dev *os.File, mount WorldMount) (WorldServer, Presentation, error) { - fd, err := unix.Dup(int(dev.Fd())) + root, err := gofs.NewLoopbackRoot(w.dir) if err != nil { return nil, Presentation{}, err } - root, err := gofs.NewLoopbackRoot(w.dir) + return w.serveRoot(root, dev, mount) +} + +func (w *loopbackWorld) serveRoot(root gofs.InodeEmbedder, dev *os.File, mount WorldMount) (WorldServer, Presentation, error) { + fd, err := unix.Dup(int(dev.Fd())) if err != nil { - unix.Close(fd) return nil, Presentation{}, err } srv, err := fuse.NewServer(gofs.NewNodeFS(root, &gofs.Options{}), fmt.Sprintf("/dev/fd/%d", fd), &fuse.MountOptions{}) @@ -367,6 +413,48 @@ func (w *loopbackWorld) Stop() error { } } +// hangWorld is a loopbackWorld in which writes to a file named hang get no answer, even when the writer is killed, and Stop never returns, until release closes. +type hangWorld struct { + loopbackWorld + hung, release chan struct{} // hung closes at the first write to hang + hungOnce sync.Once +} + +func (w *hangWorld) serve(_ context.Context, dev *os.File, mount WorldMount) (WorldServer, Presentation, error) { + root, err := gofs.NewLoopbackRoot(w.dir) + if err != nil { + return nil, Presentation{}, err + } + root.(*gofs.LoopbackNode).RootData.NewNode = func(r *gofs.LoopbackRoot, _ *gofs.Inode, name string, _ *syscall.Stat_t) gofs.InodeEmbedder { + n := &gofs.LoopbackNode{RootData: r} + if name == "hang" { + return &hangNode{LoopbackNode: n, w: w} + } + return n + } + _, p, err := w.serveRoot(root, dev, mount) + if err != nil { + return nil, p, err + } + return w, p, nil +} + +func (w *hangWorld) Stop() error { + <-w.release + return w.loopbackWorld.Stop() +} + +type hangNode struct { + *gofs.LoopbackNode + w *hangWorld +} + +func (n *hangNode) Write(context.Context, gofs.FileHandle, []byte, int64) (uint32, syscall.Errno) { + n.w.hungOnce.Do(func() { close(n.w.hung) }) + <-n.w.release + return 0, syscall.EIO +} + // serveBroker listens in the view's network namespace, as the broker does. func serveBroker(t *testing.T) func(*os.File) error { return func(netns *os.File) error { @@ -475,6 +563,14 @@ func runHelper(mode string) int { return 0 case "noop": return 0 + case "hang": + f, err := os.OpenFile("/data/hang", os.O_WRONLY, 0) + if err != nil { + fmt.Fprintln(os.Stderr, err) + return 1 + } + f.Write([]byte("never answered")) + return 0 } fmt.Fprintf(os.Stderr, "unknown helper %q\n", mode) return 2 @@ -571,12 +667,31 @@ var viewChecks = []struct { } var pids []int for _, e := range entries { - if pid, err := strconv.Atoi(e.Name()); err == nil { + if pid, err := strconv.Atoi(e.Name()); err == nil && pid != relayPID() { pids = append(pids, pid) } } if !slices.Equal(pids, []int{1, os.Getpid()}) { - return fmt.Errorf("pids %v", pids) + return fmt.Errorf("pids %v besides the relay", pids) + } + return nil + }}, + {"the relay runs as the view user without privileges", func() error { + pid := relayPID() + if pid == 0 { + return errors.New("no relay") + } + id := fmt.Sprintf("%d\t%[1]d\t%[1]d\t%[1]d", viewID) + return statusHas(fmt.Sprintf("/proc/%d/status", pid), map[string]string{"Uid": id, "Gid": id, "CapPrm": noCaps, "CapEff": noCaps, "CapBnd": noCaps, "NoNewPrivs": "1"}) + }}, + {"the relay's socket accepts the view user and cannot be replaced", func() error { + c, err := net.Dial("unix", processshim.SocketPath) + if err != nil { + return err + } + c.Close() + if err := os.Remove(processshim.SocketPath); !errors.Is(err, syscall.EROFS) { + return fmt.Errorf("remove: %v, want EROFS", err) } return nil }}, @@ -617,6 +732,19 @@ var viewChecks = []struct { } // onlyStdio finds inherited fds. Package initialisers in this binary open fds of their own, but Go opens every fd close-on-exec, so an fd without FD_CLOEXEC crossed the exec. +// relayPID returns the pid of the view's relay, or 0. +func relayPID() int { + want := []byte(strings.Join(processshim.RelayArgs, "\x00") + "\x00") + cmdlines, _ := filepath.Glob("/proc/[0-9]*/cmdline") + for _, p := range cmdlines { + if b, err := os.ReadFile(p); err == nil && bytes.Equal(b, want) { + pid, _ := strconv.Atoi(filepath.Base(filepath.Dir(p))) + return pid + } + } + return 0 +} + func onlyStdio() error { entries, err := os.ReadDir("/proc/self/fd") if err != nil { @@ -639,12 +767,18 @@ func onlyStdio() error { return nil } +const noCaps = "0000000000000000" + func noPrivileges() error { - status, err := os.ReadFile("/proc/self/status") + return statusHas("/proc/self/status", map[string]string{"CapInh": noCaps, "CapPrm": noCaps, "CapEff": noCaps, "CapBnd": noCaps, "CapAmb": noCaps, "NoNewPrivs": "1"}) +} + +// statusHas checks fields of a /proc//status file. +func statusHas(path string, want map[string]string) error { + status, err := os.ReadFile(path) if err != nil { return err } - want := map[string]string{"CapInh": "0000000000000000", "CapPrm": "0000000000000000", "CapEff": "0000000000000000", "CapBnd": "0000000000000000", "CapAmb": "0000000000000000", "NoNewPrivs": "1"} for _, line := range strings.Split(string(status), "\n") { k, v, _ := strings.Cut(line, ":") if w, ok := want[k]; ok { diff --git a/apps/daemon/internal/sessionview/view_other.go b/apps/daemon/internal/sessionview/view_other.go index 8770549d..8861e1ad 100644 --- a/apps/daemon/internal/sessionview/view_other.go +++ b/apps/daemon/internal/sessionview/view_other.go @@ -27,3 +27,4 @@ func (*View) Close() error { return nil } func (*View) Stdin() *os.File { return nil } func (*View) Stdout() *os.File { return nil } func (*View) Stderr() *os.File { return nil } +func (*View) Relay() *os.File { return nil } diff --git a/apps/daemon/internal/worldfs/world_linux_test.go b/apps/daemon/internal/worldfs/world_linux_test.go index a2a5d6af..2bfdf0da 100644 --- a/apps/daemon/internal/worldfs/world_linux_test.go +++ b/apps/daemon/internal/worldfs/world_linux_test.go @@ -16,6 +16,7 @@ import ( "testing" "time" + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/processshim" "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/sessionview" "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/worldfs" "github.com/MiniMax-AI/OpenAgentCore/apps/sandboxio/fileservicetest" @@ -35,9 +36,12 @@ const ( shimMarker = "oac-test-shim" ) -// The test binary is also the launcher, the Harness and the shim inside the view. +// The test binary is also the launcher, the Harness, the shim and the relay inside the view. func TestMain(m *testing.M) { sessionview.Init() + if processshim.Relaying() { + os.Exit(processshim.Relay()) + } if filepath.Base(os.Args[0]) == "sh" { fmt.Println(shimMarker) os.Exit(0) diff --git a/apps/sandboxio/internal/processservice/service.go b/apps/sandboxio/internal/processservice/service.go index f0c525b9..7b69a8e3 100644 --- a/apps/sandboxio/internal/processservice/service.go +++ b/apps/sandboxio/internal/processservice/service.go @@ -85,7 +85,7 @@ func New(cfg Config) (*Service, error) { IOModes: []sp.IOMode{sp.IOPipes, sp.IOPTY}, Signals: signals, SignalTargets: []sp.SignalTarget{sp.TargetLeader, sp.TargetInitialProcessGroup, sp.TargetPTYForegroundGroup, sp.TargetScope}, - PTYModes: supportedModes(), + PTYModes: sp.LinuxModes(), MaxStartBytes: sandboxwire.MaxPayload, MaxDataBytes: sandboxwire.MaxChunk, MaxActiveOperations: uint32(cfg.MaxActiveOperations), diff --git a/apps/sandboxio/internal/processservice/termios.go b/apps/sandboxio/internal/processservice/termios.go index 7cb1a010..41d8fbee 100644 --- a/apps/sandboxio/internal/processservice/termios.go +++ b/apps/sandboxio/internal/processservice/termios.go @@ -4,137 +4,23 @@ package processservice import ( "os" - "slices" "golang.org/x/sys/unix" sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" ) -type termField uint8 - -const ( - fieldCC termField = iota - fieldIflag - fieldLflag - fieldOflag - fieldCsize - fieldCflag -) - -// termMode maps a supported terminal mode to termios: a control-character -// index, a flag bit, or a character size. -type termMode struct { - field termField - value uint32 -} - -// termModes are the RFC 4254 modes Linux implements. VDSUSP, VFLUSH, VSWTCH -// and VSTATUS have no Linux character, and a PTY has no line speed. -var termModes = map[sp.PTYMode]termMode{ - sp.ModeVINTR: {fieldCC, unix.VINTR}, - sp.ModeVQUIT: {fieldCC, unix.VQUIT}, - sp.ModeVERASE: {fieldCC, unix.VERASE}, - sp.ModeVKILL: {fieldCC, unix.VKILL}, - sp.ModeVEOF: {fieldCC, unix.VEOF}, - sp.ModeVEOL: {fieldCC, unix.VEOL}, - sp.ModeVEOL2: {fieldCC, unix.VEOL2}, - sp.ModeVSTART: {fieldCC, unix.VSTART}, - sp.ModeVSTOP: {fieldCC, unix.VSTOP}, - sp.ModeVSUSP: {fieldCC, unix.VSUSP}, - sp.ModeVREPRINT: {fieldCC, unix.VREPRINT}, - sp.ModeVWERASE: {fieldCC, unix.VWERASE}, - sp.ModeVLNEXT: {fieldCC, unix.VLNEXT}, - sp.ModeVDISCARD: {fieldCC, unix.VDISCARD}, - sp.ModeIGNPAR: {fieldIflag, unix.IGNPAR}, - sp.ModePARMRK: {fieldIflag, unix.PARMRK}, - sp.ModeINPCK: {fieldIflag, unix.INPCK}, - sp.ModeISTRIP: {fieldIflag, unix.ISTRIP}, - sp.ModeINLCR: {fieldIflag, unix.INLCR}, - sp.ModeIGNCR: {fieldIflag, unix.IGNCR}, - sp.ModeICRNL: {fieldIflag, unix.ICRNL}, - sp.ModeIUCLC: {fieldIflag, unix.IUCLC}, - sp.ModeIXON: {fieldIflag, unix.IXON}, - sp.ModeIXANY: {fieldIflag, unix.IXANY}, - sp.ModeIXOFF: {fieldIflag, unix.IXOFF}, - sp.ModeIMAXBEL: {fieldIflag, unix.IMAXBEL}, - sp.ModeIUTF8: {fieldIflag, unix.IUTF8}, - sp.ModeISIG: {fieldLflag, unix.ISIG}, - sp.ModeICANON: {fieldLflag, unix.ICANON}, - sp.ModeXCASE: {fieldLflag, unix.XCASE}, - sp.ModeECHO: {fieldLflag, unix.ECHO}, - sp.ModeECHOE: {fieldLflag, unix.ECHOE}, - sp.ModeECHOK: {fieldLflag, unix.ECHOK}, - sp.ModeECHONL: {fieldLflag, unix.ECHONL}, - sp.ModeNOFLSH: {fieldLflag, unix.NOFLSH}, - sp.ModeTOSTOP: {fieldLflag, unix.TOSTOP}, - sp.ModeIEXTEN: {fieldLflag, unix.IEXTEN}, - sp.ModeECHOCTL: {fieldLflag, unix.ECHOCTL}, - sp.ModeECHOKE: {fieldLflag, unix.ECHOKE}, - sp.ModePENDIN: {fieldLflag, unix.PENDIN}, - sp.ModeOPOST: {fieldOflag, unix.OPOST}, - sp.ModeOLCUC: {fieldOflag, unix.OLCUC}, - sp.ModeONLCR: {fieldOflag, unix.ONLCR}, - sp.ModeOCRNL: {fieldOflag, unix.OCRNL}, - sp.ModeONOCR: {fieldOflag, unix.ONOCR}, - sp.ModeONLRET: {fieldOflag, unix.ONLRET}, - sp.ModeCS7: {fieldCsize, unix.CS7}, - sp.ModeCS8: {fieldCsize, unix.CS8}, - sp.ModePARENB: {fieldCflag, unix.PARENB}, - sp.ModePARODD: {fieldCflag, unix.PARODD}, -} - -func supportedModes() []sp.PTYMode { - modes := make([]sp.PTYMode, 0, len(termModes)) - for m := range termModes { - modes = append(modes, m) - } - slices.Sort(modes) - return modes -} - -// applyModes sets the requested modes on the terminal. A character value of -// 255 disables the character. CS7 or CS8 set to 1 selects that character -// size; set to 0 it leaves the size unchanged. +// applyModes sets the requested modes on the terminal. func applyModes(tty *os.File, modes []sp.PTYModeValue) error { fd := int(tty.Fd()) t, err := unix.IoctlGetTermios(fd, unix.TCGETS) if err != nil { return err } - for _, mv := range modes { - m := termModes[mv.Mode] - switch m.field { - case fieldCC: - c := uint8(mv.Value) - if mv.Value == sp.DisabledChar { - c = 0 // _POSIX_VDISABLE - } - t.Cc[m.value] = c - case fieldIflag: - t.Iflag = setFlag(t.Iflag, m.value, mv.Value) - case fieldLflag: - t.Lflag = setFlag(t.Lflag, m.value, mv.Value) - case fieldOflag: - t.Oflag = setFlag(t.Oflag, m.value, mv.Value) - case fieldCflag: - t.Cflag = setFlag(t.Cflag, m.value, mv.Value) - case fieldCsize: - if mv.Value == 1 { - t.Cflag = t.Cflag&^unix.CSIZE | m.value - } - } - } + sp.ApplyModes(t, modes) return unix.IoctlSetTermios(fd, unix.TCSETS, t) } -func setFlag(flags, bit, value uint32) uint32 { - if value == 1 { - return flags | bit - } - return flags &^ bit -} - func winsize(s sp.WindowSize) *unix.Winsize { return &unix.Winsize{Row: s.Rows, Col: s.Cols, Xpixel: s.XPixels, Ypixel: s.YPixels} } diff --git a/apps/sandboxio/testdata/processserve/main.go b/apps/sandboxio/testdata/processserve/main.go new file mode 100644 index 00000000..b5ebf40f --- /dev/null +++ b/apps/sandboxio/testdata/processserve/main.go @@ -0,0 +1,56 @@ +// Command processserve serves the Linux process service on a Unix socket for +// tests of its callers, such as apps/daemon/internal/processbroker. Every +// connection is one stream of the same attachment, so a caller that redials +// finds its operations again, as it would through a Link. +// +// Usage: processserve . It prints "ready" once it listens. +package main + +import ( + "context" + "fmt" + "net" + "os" + + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/sandboxio/internal/processservice" + sp "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxprocess" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +func main() { + processservice.Init() + if err := run(); err != nil { + fmt.Fprintln(os.Stderr, "processserve:", err) + os.Exit(1) + } +} + +func run() error { + if len(os.Args) != 2 { + return fmt.Errorf("usage: processserve ") + } + if err := unix.Prctl(unix.PR_SET_CHILD_SUBREAPER, 1, 0, 0, 0); err != nil { + return err + } + ctx := context.Background() + go processservice.Reap(ctx) + svc, err := processservice.New(processservice.DefaultConfig()) + if err != nil { + return err + } + ln, err := net.Listen("unix", os.Args[1]) + if err != nil { + return err + } + att := sp.Attachment{ID: sandboxwire.NewID()} + fmt.Println("ready") + for { + c, err := ln.Accept() + if err != nil { + return err + } + go sp.Serve(ctx, c, att, svc) + } +} diff --git a/contracts/agents-api/harness-onboarding.md b/contracts/agents-api/harness-onboarding.md index 22a44331..b35ecc14 100644 --- a/contracts/agents-api/harness-onboarding.md +++ b/contracts/agents-api/harness-onboarding.md @@ -277,10 +277,10 @@ An agent host runs the Harness outside the sandbox, in a per-Session view. The v `View.Validate` checks the declaration without touching the host: - view and host paths are absolute and clean; -- closure names are single path components other than `bin`, `home` and `run`, which the agent host uses for the shims, the Session home and the process broker; +- closure names are single path components other than `bin`, `home` and `run`, which the agent host uses for the shims, the Session home and the process relay; - shim paths, overlays and masks do not overlap each other or `/`, and stay out of the trees the view builds itself: `/.oac`, `/proc` and `/dev` (`ViewReserved`); - each `LocalExec` entry lies in a closure directory or an `Exec` overlay; -- shim names and `ForwardEnv` names are unique, a variable name contains no `=`, and `ForwardEnv` names no variable the view or the broker sets ([Environment](#environment)); +- shim names and `ForwardEnv` names are unique, no shim is named `oac-process-shim`, which is the process relay's, a variable name contains no `=`, and `ForwardEnv` names no variable the view or the broker sets ([Environment](#environment)); - `Proxy` is one of the two values and `Executor` is non-nil. `harness.go` defines the view layout once, and `sessionview` builds views from it. The agent host checks its own overlays, such as `/etc/passwd`, against the declaration when it builds the view. diff --git a/internal/sandboxprocess/termios_linux.go b/internal/sandboxprocess/termios_linux.go new file mode 100644 index 00000000..6938da50 --- /dev/null +++ b/internal/sandboxprocess/termios_linux.go @@ -0,0 +1,175 @@ +//go:build linux + +package sandboxprocess + +import ( + "slices" + + "golang.org/x/sys/unix" +) + +// This file is the Linux termios projection of the terminal modes, used both +// to read a local terminal into a PTYSpec and to apply a PTYSpec to a new one. + +type termField uint8 + +const ( + fieldCC termField = iota + fieldIflag + fieldLflag + fieldOflag + fieldCsize + fieldCflag +) + +// termMode maps a mode to termios: a control-character index, a flag bit, or +// a character size. +type termMode struct { + field termField + value uint32 +} + +// termModes are the modes Linux implements. VDSUSP, VFLUSH, VSWTCH and VSTATUS +// have no Linux character, and a PTY has no line speed. +var termModes = map[PTYMode]termMode{ + ModeVINTR: {fieldCC, unix.VINTR}, + ModeVQUIT: {fieldCC, unix.VQUIT}, + ModeVERASE: {fieldCC, unix.VERASE}, + ModeVKILL: {fieldCC, unix.VKILL}, + ModeVEOF: {fieldCC, unix.VEOF}, + ModeVEOL: {fieldCC, unix.VEOL}, + ModeVEOL2: {fieldCC, unix.VEOL2}, + ModeVSTART: {fieldCC, unix.VSTART}, + ModeVSTOP: {fieldCC, unix.VSTOP}, + ModeVSUSP: {fieldCC, unix.VSUSP}, + ModeVREPRINT: {fieldCC, unix.VREPRINT}, + ModeVWERASE: {fieldCC, unix.VWERASE}, + ModeVLNEXT: {fieldCC, unix.VLNEXT}, + ModeVDISCARD: {fieldCC, unix.VDISCARD}, + ModeIGNPAR: {fieldIflag, unix.IGNPAR}, + ModePARMRK: {fieldIflag, unix.PARMRK}, + ModeINPCK: {fieldIflag, unix.INPCK}, + ModeISTRIP: {fieldIflag, unix.ISTRIP}, + ModeINLCR: {fieldIflag, unix.INLCR}, + ModeIGNCR: {fieldIflag, unix.IGNCR}, + ModeICRNL: {fieldIflag, unix.ICRNL}, + ModeIUCLC: {fieldIflag, unix.IUCLC}, + ModeIXON: {fieldIflag, unix.IXON}, + ModeIXANY: {fieldIflag, unix.IXANY}, + ModeIXOFF: {fieldIflag, unix.IXOFF}, + ModeIMAXBEL: {fieldIflag, unix.IMAXBEL}, + ModeIUTF8: {fieldIflag, unix.IUTF8}, + ModeISIG: {fieldLflag, unix.ISIG}, + ModeICANON: {fieldLflag, unix.ICANON}, + ModeXCASE: {fieldLflag, unix.XCASE}, + ModeECHO: {fieldLflag, unix.ECHO}, + ModeECHOE: {fieldLflag, unix.ECHOE}, + ModeECHOK: {fieldLflag, unix.ECHOK}, + ModeECHONL: {fieldLflag, unix.ECHONL}, + ModeNOFLSH: {fieldLflag, unix.NOFLSH}, + ModeTOSTOP: {fieldLflag, unix.TOSTOP}, + ModeIEXTEN: {fieldLflag, unix.IEXTEN}, + ModeECHOCTL: {fieldLflag, unix.ECHOCTL}, + ModeECHOKE: {fieldLflag, unix.ECHOKE}, + ModePENDIN: {fieldLflag, unix.PENDIN}, + ModeOPOST: {fieldOflag, unix.OPOST}, + ModeOLCUC: {fieldOflag, unix.OLCUC}, + ModeONLCR: {fieldOflag, unix.ONLCR}, + ModeOCRNL: {fieldOflag, unix.OCRNL}, + ModeONOCR: {fieldOflag, unix.ONOCR}, + ModeONLRET: {fieldOflag, unix.ONLRET}, + ModeCS7: {fieldCsize, unix.CS7}, + ModeCS8: {fieldCsize, unix.CS8}, + ModePARENB: {fieldCflag, unix.PARENB}, + ModePARODD: {fieldCflag, unix.PARODD}, +} + +// LinuxModes returns the modes Linux termios implements, sorted. +func LinuxModes() []PTYMode { + modes := make([]PTYMode, 0, len(termModes)) + for m := range termModes { + modes = append(modes, m) + } + slices.Sort(modes) + return modes +} + +// ReadModes returns the value in t of each mode in modes that Linux +// implements. A disabled control character reads as DisabledChar. +func ReadModes(t *unix.Termios, modes []PTYMode) []PTYModeValue { + var out []PTYModeValue + for _, mode := range modes { + m, ok := termModes[mode] + if !ok { + continue + } + var v uint32 + switch m.field { + case fieldCC: + v = uint32(t.Cc[m.value]) + if v == 0 { // _POSIX_VDISABLE + v = DisabledChar + } + case fieldIflag: + v = flagValue(t.Iflag, m.value) + case fieldLflag: + v = flagValue(t.Lflag, m.value) + case fieldOflag: + v = flagValue(t.Oflag, m.value) + case fieldCflag: + v = flagValue(t.Cflag, m.value) + case fieldCsize: + if t.Cflag&unix.CSIZE == m.value { + v = 1 + } + } + out = append(out, PTYModeValue{Mode: mode, Value: v}) + } + return out +} + +// ApplyModes sets modes on t, skipping modes Linux does not implement. A +// character value of DisabledChar disables the character. CS7 or CS8 set to 1 +// selects that character size; set to 0 it leaves the size unchanged. +func ApplyModes(t *unix.Termios, modes []PTYModeValue) { + for _, mv := range modes { + m, ok := termModes[mv.Mode] + if !ok { + continue + } + switch m.field { + case fieldCC: + c := uint8(mv.Value) + if mv.Value == DisabledChar { + c = 0 // _POSIX_VDISABLE + } + t.Cc[m.value] = c + case fieldIflag: + t.Iflag = setFlag(t.Iflag, m.value, mv.Value) + case fieldLflag: + t.Lflag = setFlag(t.Lflag, m.value, mv.Value) + case fieldOflag: + t.Oflag = setFlag(t.Oflag, m.value, mv.Value) + case fieldCflag: + t.Cflag = setFlag(t.Cflag, m.value, mv.Value) + case fieldCsize: + if mv.Value == 1 { + t.Cflag = t.Cflag&^unix.CSIZE | m.value + } + } + } +} + +func flagValue(flags, bit uint32) uint32 { + if flags&bit != 0 { + return 1 + } + return 0 +} + +func setFlag(flags, bit, value uint32) uint32 { + if value == 1 { + return flags | bit + } + return flags &^ bit +}