335 lines
8.2 KiB
Go
335 lines
8.2 KiB
Go
// toolsrv — 9P tool execution server.
|
|
//
|
|
// Executes tool calls in a sandboxed environment. Used both locally
|
|
// (spawned by olliesrv) and remotely (deployed via SSH bootstrap).
|
|
//
|
|
// Usage:
|
|
//
|
|
// toolsrv serve --cwd /path/to/project --listen /path/to/sock [--yolo]
|
|
//
|
|
// Listens on a Unix socket and serves a 9P filesystem for tool execution.
|
|
// Authentication happens via Tauth: first client to connect sets the secret,
|
|
// subsequent clients must provide the same secret.
|
|
package main
|
|
|
|
import (
|
|
"context"
|
|
"flag"
|
|
"fmt"
|
|
"net"
|
|
"os"
|
|
"os/signal"
|
|
"os/user"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"syscall"
|
|
"time"
|
|
|
|
"9fans.net/go/plan9"
|
|
p9client "9fans.net/go/plan9/client"
|
|
|
|
"ollie/cmd/toolsrv/internal/registry"
|
|
"ollie/cmd/toolsrv/internal/sandbox"
|
|
"ollie/cmd/toolsrv/internal/server"
|
|
|
|
olog "ollie/log"
|
|
"ollie/util"
|
|
)
|
|
|
|
var (
|
|
cwd = flag.String("cwd", ".", "working directory for execution")
|
|
listenPath = flag.String("listen", "", "Unix socket path to listen on (required)")
|
|
sessionID = flag.String("session-id", "", "session ID for tool registry")
|
|
yolo = flag.Bool("yolo", false, "skip sandbox enforcement")
|
|
noAuth = flag.Bool("no-auth", false, "disable Tauth authentication (for debugging)")
|
|
idleTimeout = flag.Duration("idle-timeout", 30*time.Second, "exit after this duration with no connections (0 to disable)")
|
|
)
|
|
|
|
// logger is the toolsrv logger.
|
|
var logger *olog.Logger
|
|
|
|
// activeConns tracks the number of active 9P connections.
|
|
var activeConns atomic.Int64
|
|
|
|
// hadConnection is set to true once the first client connects.
|
|
var hadConnection atomic.Bool
|
|
|
|
// procInterruptWG tracks completion notifications during shutdown.
|
|
var procInterruptWG sync.WaitGroup
|
|
|
|
func main() {
|
|
if len(os.Args) > 1 && os.Args[1] == "sandbox-exec" {
|
|
if err := sandbox.ExecHelper(os.Args[2:]); err != nil {
|
|
fmt.Fprintln(os.Stderr, err)
|
|
os.Exit(1)
|
|
}
|
|
return
|
|
}
|
|
if len(os.Args) < 2 || os.Args[1] != "serve" {
|
|
fmt.Fprintln(os.Stderr, "usage: toolsrv serve --cwd <path> --listen <socket> [--yolo]")
|
|
os.Exit(1)
|
|
}
|
|
flag.CommandLine.Parse(os.Args[2:])
|
|
|
|
if *listenPath == "" {
|
|
fmt.Fprintln(os.Stderr, "error: --listen is required")
|
|
os.Exit(1)
|
|
}
|
|
|
|
// Expand ~ in CWD (shell doesn't expand ~ in arguments).
|
|
*cwd = util.ExpandHome(*cwd)
|
|
|
|
// Set up logger
|
|
sink := olog.NewSink(os.Stdout, os.Stderr, olog.LevelInfo)
|
|
defer sink.Flush()
|
|
logger = sink.NewLogger("toolsrv")
|
|
|
|
util.EnsureEnv()
|
|
|
|
// Native Landlock is applied by the sandbox-exec child; no external wrapper is required.
|
|
|
|
// Create tool registry
|
|
toolReg, err := registry.New()
|
|
if err != nil {
|
|
logger.Warn("tool registry: %v", err)
|
|
}
|
|
|
|
// Create server state
|
|
srv := server.NewServer()
|
|
srv.Fs.SetCWD(*cwd)
|
|
srv.SetYolo(*yolo)
|
|
if toolReg != nil {
|
|
srv.SetRegistry(toolReg)
|
|
}
|
|
|
|
// Set up signal handling
|
|
runCtx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
sigCh := make(chan os.Signal, 1)
|
|
signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
|
|
defer signal.Stop(sigCh)
|
|
go func() {
|
|
<-sigCh
|
|
cancel()
|
|
}()
|
|
|
|
// Set up callback to push completed proc output to agent prompt
|
|
srv.Fs.OnProcExit = func(proc *server.Proc) {
|
|
if proc.SessionID == "" || proc.AgentID == "" {
|
|
return
|
|
}
|
|
procInterruptWG.Add(1)
|
|
go func() {
|
|
defer procInterruptWG.Done()
|
|
pushProcInterrupt(runCtx, proc)
|
|
}()
|
|
}
|
|
|
|
tree := srv.BuildTree()
|
|
if tree == nil {
|
|
logger.Error("failed to build filesystem tree")
|
|
os.Exit(1)
|
|
}
|
|
|
|
// Kill all procs on shutdown
|
|
defer procInterruptWG.Wait()
|
|
defer srv.Fs.KillAll()
|
|
|
|
// Start idle timeout monitor
|
|
if *idleTimeout > 0 {
|
|
go idleMonitor(runCtx, cancel, *idleTimeout)
|
|
}
|
|
|
|
// Start proc GC
|
|
srv.Fs.StartProcGC(runCtx)
|
|
|
|
// Remove stale socket if it exists
|
|
os.Remove(*listenPath)
|
|
|
|
// Listen on Unix socket
|
|
ln, err := net.Listen("unix", *listenPath)
|
|
if err != nil {
|
|
logger.Error("listen failed: %v", err)
|
|
os.Exit(1)
|
|
}
|
|
defer ln.Close()
|
|
defer os.Remove(*listenPath)
|
|
|
|
logger.Info("listening on %s", *listenPath)
|
|
|
|
// Accept connections
|
|
go func() {
|
|
<-runCtx.Done()
|
|
ln.Close()
|
|
}()
|
|
|
|
for {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
select {
|
|
case <-runCtx.Done():
|
|
return
|
|
default:
|
|
logger.Warn("accept error: %v", err)
|
|
continue
|
|
}
|
|
}
|
|
activeConns.Add(1)
|
|
hadConnection.Store(true)
|
|
go func() {
|
|
defer activeConns.Add(-1)
|
|
serve9P(runCtx, conn, tree, srv, *noAuth)
|
|
}()
|
|
}
|
|
}
|
|
|
|
// idleMonitor exits the process if there are no active connections for the
|
|
// given duration. It waits for the first connection before monitoring (with
|
|
// a startup grace period of 60s).
|
|
func idleMonitor(ctx context.Context, cancel context.CancelFunc, timeout time.Duration) {
|
|
const startupGrace = 60 * time.Second
|
|
|
|
// Wait for first connection (or startup grace period)
|
|
deadline := time.NewTimer(startupGrace)
|
|
defer deadline.Stop()
|
|
poll := time.NewTicker(1 * time.Second)
|
|
defer poll.Stop()
|
|
for {
|
|
select {
|
|
case <-ctx.Done():
|
|
return
|
|
case <-deadline.C:
|
|
if !hadConnection.Load() {
|
|
logger.Info("no connections within startup grace period, exiting")
|
|
cancel()
|
|
return
|
|
}
|
|
goto monitor
|
|
case <-poll.C:
|
|
if hadConnection.Load() {
|
|
goto monitor
|
|
}
|
|
}
|
|
}
|
|
|
|
monitor:
|
|
// Monitor for idle (zero connections after having had at least one)
|
|
idleSince := time.Time{}
|
|
poll.Reset(5 * time.Second)
|
|
for {
|
|
select {
|
|
case <-ctx.Done():
|
|
return
|
|
case <-poll.C:
|
|
if activeConns.Load() > 0 {
|
|
idleSince = time.Time{}
|
|
poll.Reset(5 * time.Second)
|
|
continue
|
|
}
|
|
// No active connections
|
|
if idleSince.IsZero() {
|
|
idleSince = time.Now()
|
|
logger.Info("all connections closed, idle timer started (timeout=%v)", timeout)
|
|
} else if time.Since(idleSince) >= timeout {
|
|
logger.Info("idle timeout reached, exiting")
|
|
cancel()
|
|
return
|
|
}
|
|
poll.Reset(5 * time.Second)
|
|
}
|
|
}
|
|
}
|
|
|
|
// pushProcInterrupt pushes the completed proc output to the agent's prompt
|
|
// and publishes a proc.exit event.
|
|
func pushProcInterrupt(ctx context.Context, proc *server.Proc) {
|
|
// Build the interrupt message
|
|
proc.Lock()
|
|
output := proc.OutputString()
|
|
exitCode := proc.ExitCode
|
|
cmd := proc.Cmd
|
|
tool := proc.Tool
|
|
pid := proc.ID
|
|
sid := proc.SessionID
|
|
aid := proc.AgentID
|
|
proc.Unlock()
|
|
|
|
msg := fmt.Sprintf("<system-proc-complete id=\"%d\" cmd=\"%s\" exit=\"%d\">\n%s\n</system-proc-complete>",
|
|
pid, escapeAttr(cmd), exitCode, output)
|
|
|
|
// Connect to olliesrv
|
|
u, err := user.Current()
|
|
if err != nil {
|
|
logger.Warn("pushProcInterrupt: failed to get current user: %v", err)
|
|
return
|
|
}
|
|
display := os.Getenv("DISPLAY")
|
|
if display == "" {
|
|
display = ":0"
|
|
}
|
|
sockPath := fmt.Sprintf("/tmp/ns.%s.%s/ollie", u.Username, display)
|
|
|
|
ctx, cancel := context.WithTimeout(ctx, 10*time.Second)
|
|
defer cancel()
|
|
conn, err := dialUnixContext(ctx, sockPath)
|
|
if err != nil {
|
|
logger.Warn("pushProcInterrupt: failed to connect to olliesrv: %v", err)
|
|
return
|
|
}
|
|
p9conn, err := p9client.NewConn(conn)
|
|
if err != nil {
|
|
conn.Close()
|
|
logger.Warn("pushProcInterrupt: failed to initialize 9P connection: %v", err)
|
|
return
|
|
}
|
|
defer p9conn.Close()
|
|
fsys, err := p9conn.Attach(nil, u.Username, "")
|
|
if err != nil {
|
|
conn.Close()
|
|
logger.Warn("pushProcInterrupt: failed to attach to olliesrv: %v", err)
|
|
return
|
|
}
|
|
defer fsys.Close()
|
|
|
|
// Publish proc.exit event
|
|
eventPath := "event.pub"
|
|
eventTopic := fmt.Sprintf("session.%s.agent.%s.proc.exit", sid, aid)
|
|
eventPayload := fmt.Sprintf("%d\t%d\t%s\t%s", pid, exitCode, tool, cmd)
|
|
if evFid, err := fsys.Open(eventPath, plan9.OWRITE); err == nil {
|
|
evFid.Write([]byte(eventTopic + "\t" + eventPayload))
|
|
evFid.Close()
|
|
}
|
|
|
|
// Write to session/{sid}/agent/{aid}/prompt
|
|
promptPath := fmt.Sprintf("session/%s/agent/%s/prompt", sid, aid)
|
|
fid, err := fsys.Open(promptPath, plan9.OWRITE)
|
|
if err != nil {
|
|
logger.Warn("pushProcInterrupt: failed to open prompt: %v", err)
|
|
return
|
|
}
|
|
_, err = fid.Write([]byte(msg))
|
|
closeErr := fid.Close()
|
|
if err == nil {
|
|
err = closeErr
|
|
}
|
|
if err != nil {
|
|
logger.Warn("pushProcInterrupt: failed to write prompt: %v", err)
|
|
return
|
|
}
|
|
|
|
logger.Info("pushed proc %d interrupt to %s", pid, promptPath)
|
|
}
|
|
|
|
func dialUnixContext(ctx context.Context, path string) (net.Conn, error) {
|
|
var d net.Dialer
|
|
return d.DialContext(ctx, "unix", path)
|
|
}
|
|
|
|
func escapeAttr(s string) string {
|
|
s = strings.ReplaceAll(s, "&", "&")
|
|
s = strings.ReplaceAll(s, "\"", """)
|
|
s = strings.ReplaceAll(s, "<", "<")
|
|
return s
|
|
}
|