ollie/cmd/toolsrv/internal/fs/proc.go

632 lines
14 KiB
Go

// proc.go - Process management and tool operations for toolsrv.
package fs
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"os"
"runtime"
"strconv"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"ollie/cmd/toolsrv/internal/exec"
"ollie/cmd/toolsrv/internal/registry"
"ollie/toolsrv"
)
// State implements the filesystem state for toolsrv.
type State struct {
mu sync.RWMutex
cwd string
env map[string]string
registry *registry.Registry
sessID string
yolo bool
// Process management
procMu sync.Mutex
procs map[int]*Proc
nextID int
procLimit int
// Callbacks
OnToolsChanged func()
}
// Proc represents a running or completed process.
type Proc struct {
ID int
Tool string
Args map[string]string
StartTime time.Time
EndTime time.Time
ExitCode int
Exited bool
LastRead time.Time // last time output was read; GC TTL counts from here
mu sync.Mutex
output bytes.Buffer
done chan struct{}
cancel context.CancelFunc
proc *os.Process // for signaling
}
// NewState creates a new filesystem state.
func NewState(cwd string) *State {
return &State{
cwd: cwd,
env: make(map[string]string),
procs: make(map[int]*Proc),
nextID: 1,
procLimit: 32,
}
}
// SetRegistry configures the tool registry and session ID.
func (st *State) SetRegistry(r *registry.Registry, sessID string) {
st.mu.Lock()
st.registry = r
st.sessID = sessID
st.mu.Unlock()
}
// SetYolo enables/disables sandbox bypass.
func (st *State) SetYolo(yolo bool) {
st.mu.Lock()
st.yolo = yolo
st.mu.Unlock()
}
// CWD returns the current working directory.
func (st *State) CWD() string {
st.mu.RLock()
defer st.mu.RUnlock()
return st.cwd
}
// SetCWD sets the current working directory.
func (st *State) SetCWD(dir string) {
st.mu.Lock()
st.cwd = dir
st.mu.Unlock()
}
// SetEnv sets an environment variable.
func (st *State) SetEnv(key, value string) {
st.mu.Lock()
if st.env == nil {
st.env = make(map[string]string)
}
st.env[key] = value
st.mu.Unlock()
}
// GetEnv returns an environment variable value.
func (st *State) GetEnv(key string) string {
st.mu.RLock()
defer st.mu.RUnlock()
return st.env[key]
}
// --- Tool Registry ---
// ListTools returns loaded tools for the current session.
func (st *State) ListTools() []toolsrv.ToolInfo {
st.mu.RLock()
reg := st.registry
sid := st.sessID
st.mu.RUnlock()
if reg == nil || sid == "" {
return nil
}
return reg.Loaded(sid)
}
// LoadTool loads a tool by name into the session registry.
func (st *State) LoadTool(name string) error {
st.mu.RLock()
reg := st.registry
sid := st.sessID
st.mu.RUnlock()
if reg == nil {
return fmt.Errorf("no tool registry configured")
}
if sid == "" {
return fmt.Errorf("no session ID configured")
}
if err := reg.Load(sid, name); err != nil {
return err
}
if st.OnToolsChanged != nil {
st.OnToolsChanged()
}
return nil
}
// UnloadTool removes a tool from the session registry.
func (st *State) UnloadTool(name string) error {
st.mu.RLock()
reg := st.registry
sid := st.sessID
st.mu.RUnlock()
if reg == nil {
return fmt.Errorf("no tool registry configured")
}
if sid == "" {
return fmt.Errorf("no session ID configured")
}
if err := reg.Unload(sid, name); err != nil {
return err
}
if st.OnToolsChanged != nil {
st.OnToolsChanged()
}
return nil
}
// --- Process Management ---
// allocID allocates a new process ID.
func (st *State) allocID() int {
st.procMu.Lock()
defer st.procMu.Unlock()
pid := st.nextID
st.nextID++
return pid
}
// NewProc creates a new process and starts execution.
// If background is false, blocks until completion and returns result.
// If background is true, returns immediately with pid.
func (st *State) NewProc(ctx context.Context, payload string, background bool) (result string, pid int, err error) {
// Parse payload
args := ParsePayload(payload)
toolName := args["tool"]
if toolName == "" {
return "", 0, fmt.Errorf("missing 'tool' in payload")
}
// Check proc limit
st.procMu.Lock()
if len(st.procs) >= st.procLimit {
st.procMu.Unlock()
return "", 0, fmt.Errorf("process limit reached (%d)", st.procLimit)
}
st.procMu.Unlock()
// Look up tool
st.mu.RLock()
reg := st.registry
sid := st.sessID
cwd := st.cwd
yolo := st.yolo
envCopy := make(map[string]string)
for k, v := range st.env {
envCopy[k] = v
}
st.mu.RUnlock()
if reg == nil || sid == "" {
return "", 0, fmt.Errorf("tool registry not configured")
}
info, ok := reg.Lookup(sid, toolName)
if !ok {
return "", 0, fmt.Errorf("tool not found: %s", toolName)
}
// Create proc
procCtx, cancel := context.WithCancel(ctx)
proc := &Proc{
ID: st.allocID(),
Tool: toolName,
Args: args,
StartTime: time.Now(),
done: make(chan struct{}),
cancel: cancel,
}
st.procMu.Lock()
st.procs[proc.ID] = proc
st.procMu.Unlock()
// For background procs, stream output directly into the proc buffer.
// For foreground, use nil (local buffer in exec path).
var outputWriter io.Writer
var startedCh chan *os.Process
timeout := 30 // default for foreground
if background {
outputWriter = &procWriter{proc: proc}
startedCh = make(chan *os.Process, 1)
timeout = 0 // background procs never timeout
}
// Execute in goroutine
go func() {
defer close(proc.done)
defer cancel()
out, exitCode := st.executeTool(procCtx, info, args, cwd, envCopy, yolo, timeout, outputWriter, startedCh)
proc.mu.Lock()
if !background {
proc.output.WriteString(out)
}
proc.ExitCode = exitCode
proc.Exited = true
proc.EndTime = time.Now()
proc.mu.Unlock()
}()
// For background procs, wait for the process to start and store its reference.
if background {
if p, ok := <-startedCh; ok && p != nil {
proc.mu.Lock()
proc.proc = p
proc.mu.Unlock()
}
return fmt.Sprintf("%d", proc.ID), proc.ID, nil
}
// Wait for completion
<-proc.done
proc.mu.Lock()
result = proc.output.String()
exitCode := proc.ExitCode
proc.mu.Unlock()
// Auto-cleanup for foreground procs
st.procMu.Lock()
delete(st.procs, proc.ID)
st.procMu.Unlock()
if exitCode != 0 {
return result, proc.ID, fmt.Errorf("exit %d", exitCode)
}
return result, proc.ID, nil
}
// executeTool runs a tool and returns output + exit code.
// timeout: 0 = no timeout, >0 = seconds.
func (st *State) executeTool(ctx context.Context, info toolsrv.ToolInfo, args map[string]string, cwd string, envExtra map[string]string, yolo bool, timeout int, output io.Writer, started chan *os.Process) (string, int) {
// Convert args map to JSON for the execution path
jsonArgs := argsToJSON(args)
cfg := exec.Config{
CWD: cwd,
Env: envExtra,
Yolo: yolo,
Timeout: timeout,
Output: output,
Started: started,
}
result, err := exec.ExecuteTool(ctx, info, jsonArgs, cfg)
if err != nil {
return fmt.Sprintf("error: %v", err), 1
}
return string(result), 0
}
// GetProc returns a process by PID.
func (st *State) GetProc(pid int) *Proc {
st.procMu.Lock()
defer st.procMu.Unlock()
return st.procs[pid]
}
// ListProcs returns all process PIDs.
func (st *State) ListProcs() []int {
st.procMu.Lock()
defer st.procMu.Unlock()
pids := make([]int, 0, len(st.procs))
for pid := range st.procs {
pids = append(pids, pid)
}
return pids
}
// KillAll sends SIGKILL to all running procs and cancels their contexts.
// Called on shutdown to ensure no orphaned processes.
func (st *State) KillAll() {
st.procMu.Lock()
procs := make([]*Proc, 0, len(st.procs))
for _, p := range st.procs {
procs = append(procs, p)
}
st.procMu.Unlock()
for _, p := range procs {
p.mu.Lock()
proc := p.proc
exited := p.Exited
p.mu.Unlock()
if !exited && proc != nil {
syscall.Kill(-proc.Pid, syscall.SIGKILL)
}
p.cancel()
}
}
// DismissProc removes a process from the list.
func (st *State) DismissProc(pid int) bool {
st.procMu.Lock()
defer st.procMu.Unlock()
if _, ok := st.procs[pid]; ok {
delete(st.procs, pid)
return true
}
return false
}
const procGCTTL = 10 * time.Minute
// StartProcGC runs a background goroutine that removes exited procs
// whose output has been read and whose LastRead is older than procGCTTL.
func (st *State) StartProcGC(ctx context.Context) {
go func() {
ticker := time.NewTicker(1 * time.Minute)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
st.gcProcs()
}
}
}()
}
func (st *State) gcProcs() {
st.procMu.Lock()
defer st.procMu.Unlock()
now := time.Now()
for pid, p := range st.procs {
p.mu.Lock()
exited := p.Exited
lastRead := p.LastRead
p.mu.Unlock()
if exited && !lastRead.IsZero() && now.Sub(lastRead) >= procGCTTL {
delete(st.procs, pid)
}
}
}
// SignalProc sends a signal to a process.
func (st *State) SignalProc(pid int, sig syscall.Signal) error {
proc := st.GetProc(pid)
if proc == nil {
return fmt.Errorf("process not found: %d", pid)
}
// Signal the process group (we set Setpgid: true)
proc.mu.Lock()
p := proc.proc
proc.mu.Unlock()
if p != nil {
syscall.Kill(-p.Pid, sig)
}
// Cancel context on SIGKILL (hard stop)
if sig == syscall.SIGKILL {
proc.cancel()
}
return nil
}
// --- Host Info ---
// HostInfo returns host information.
func (st *State) HostInfo() string {
st.mu.RLock()
cwd := st.cwd
st.mu.RUnlock()
return fmt.Sprintf("platform=%s\narch=%s\ncwd=%s\ngit=%v\n",
runtime.GOOS, runtime.GOARCH, cwd, isGitRepo(cwd))
}
func isGitRepo(dir string) bool {
info, err := os.Stat(dir + "/.git")
return err == nil && info.IsDir()
}
// --- Request Handlers ---
// HandleCtl processes ctl commands.
func (st *State) HandleCtl(input string) error {
parts := strings.Fields(input)
if len(parts) == 0 {
return nil
}
switch parts[0] {
case "load":
if len(parts) < 2 {
return fmt.Errorf("load requires tool name")
}
return st.LoadTool(parts[1])
case "unload":
if len(parts) < 2 {
return fmt.Errorf("unload requires tool name")
}
return st.UnloadTool(parts[1])
case "env":
if len(parts) < 2 {
return fmt.Errorf("env requires KEY=VALUE")
}
// Rejoin in case value contains spaces
kv := strings.Join(parts[1:], " ")
idx := strings.Index(kv, "=")
if idx < 0 {
return fmt.Errorf("env requires KEY=VALUE format")
}
st.SetEnv(kv[:idx], kv[idx+1:])
return nil
case "cwd":
if len(parts) < 2 {
return fmt.Errorf("cwd requires path")
}
st.SetCWD(parts[1])
return nil
default:
return fmt.Errorf("unknown command: %s", parts[0])
}
}
// HandleToolsRead returns the tool list as JSON.
func (st *State) HandleToolsRead() string {
tools := st.ListTools()
data, _ := json.Marshal(tools)
return string(data)
}
// HandleToolsWrite loads a tool by name.
func (st *State) HandleToolsWrite(name string) error {
return st.LoadTool(strings.TrimSpace(name))
}
// HandleProcNew is the rdwr handler for /proc/new (blocking).
func (st *State) HandleProcNew(ctx context.Context, payload string) (string, error) {
result, _, err := st.NewProc(ctx, payload, false)
return result, err
}
// HandleProcNewBg is the write handler for /proc/new.bg (background).
func (st *State) HandleProcNewBg(ctx context.Context, payload string) (string, error) {
result, _, err := st.NewProc(ctx, payload, true)
return result, err
}
// HandleProcCtl processes writes to /proc/<pid>/ctl.
func (st *State) HandleProcCtl(pid int, input string) error {
parts := strings.Fields(input)
if len(parts) == 0 {
return nil
}
switch parts[0] {
case "term":
return st.SignalProc(pid, syscall.SIGTERM)
case "kill":
return st.SignalProc(pid, syscall.SIGKILL)
case "signal":
if len(parts) < 2 {
return fmt.Errorf("signal requires signal number")
}
sig, err := strconv.Atoi(parts[1])
if err != nil {
return fmt.Errorf("invalid signal: %s", parts[1])
}
return st.SignalProc(pid, syscall.Signal(sig))
case "dismiss":
st.DismissProc(pid)
return nil
default:
return fmt.Errorf("unknown command: %s", parts[0])
}
}
// --- procWriter ---
// procWriter is a thread-safe writer that streams into a Proc's output buffer.
type procWriter struct {
proc *Proc
}
func (pw *procWriter) Write(p []byte) (int, error) {
pw.proc.mu.Lock()
defer pw.proc.mu.Unlock()
return pw.proc.output.Write(p)
}
// --- Proc Methods ---
// Output returns the current output buffer and updates LastRead.
func (p *Proc) Output() string {
p.mu.Lock()
defer p.mu.Unlock()
p.LastRead = time.Now()
return p.output.String()
}
// Wait blocks until the process exits, returns exit code.
func (p *Proc) Wait() int {
<-p.done
p.mu.Lock()
defer p.mu.Unlock()
return p.ExitCode
}
// Stat returns a status string in key=value format.
func (p *Proc) Stat() string {
p.mu.Lock()
defer p.mu.Unlock()
rt := time.Since(p.StartTime)
if p.Exited {
rt = p.EndTime.Sub(p.StartTime)
return fmt.Sprintf("id=%d\nexited=true\nexit_code=%d\nruntime=%v\ntool=%s\n", p.ID, p.ExitCode, rt, p.Tool)
}
return fmt.Sprintf("id=%d\nexited=false\nruntime=%v\ntool=%s\n", p.ID, rt, p.Tool)
}
// --- Helpers ---
// argsToJSON converts the args map to JSON for the existing tool path.
func argsToJSON(args map[string]string) []byte {
var buf bytes.Buffer
buf.WriteByte('{')
first := true
for k, v := range args {
if k == "tool" {
continue // tool name is not part of the args
}
if !first {
buf.WriteByte(',')
}
first = false
buf.WriteByte('"')
buf.WriteString(escapeJSON(k))
buf.WriteString("\":")
buf.WriteByte('"')
buf.WriteString(escapeJSON(v))
buf.WriteByte('"')
}
buf.WriteByte('}')
return buf.Bytes()
}
func escapeJSON(s string) string {
s = strings.ReplaceAll(s, "\\", "\\\\")
s = strings.ReplaceAll(s, "\"", "\\\"")
s = strings.ReplaceAll(s, "\n", "\\n")
s = strings.ReplaceAll(s, "\r", "\\r")
s = strings.ReplaceAll(s, "\t", "\\t")
return s
}
// --- Atomic counters for unique IDs ---
var globalProcCounter atomic.Int64
func init() {
globalProcCounter.Store(time.Now().UnixNano() % 10000)
}