ollie/cmd/toolsrv/main.go

320 lines
7.7 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, _ := registry.New()
// Create server state
srv := server.NewServer()
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.
func pushProcInterrupt(ctx context.Context, proc *server.Proc) {
// Build the interrupt message
proc.Lock()
output := proc.OutputString()
exitCode := proc.ExitCode
cmd := proc.Cmd
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()
// 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, "&", "&amp;")
s = strings.ReplaceAll(s, "\"", "&quot;")
s = strings.ReplaceAll(s, "<", "&lt;")
return s
}