toolsrv: remove JSON-RPC, add compat stubs for 9P transition
BREAKING CHANGE: Removes all JSON-RPC toolsrv code. Deleted files: - server.go, rpcserver.go, rpcwire.go (JSON-RPC server) - conn.go, dial.go (JSON-RPC client) - exec.go, stream.go, detach.go (old execution tied to Server) - spawn.go, transport.go, processkeeper.go, remote.go (spawning) - shell_validate.go (tied to old Server) New/modified: - compat.go: Stub types to keep agent/session/fs packages compiling - Conn, Process, ProcessKeeper with TODO implementations - Spawn, SpawnRemote stubs - HostInfo, RemoteConfig, Option types - All marked TODO for 9P implementation - cmd/toolsrv/main.go: Updated to use 9P model - Requires TOOLSRV_SECRET env var - Builds fsedsl tree - serve9P placeholder for 9P protocol handling - exec9p.go: Added ToolResult types, limitedWriter, plan9Namespace The codebase compiles but toolsrv is non-functional until: 1. serve9P implements 9P protocol 2. Conn stubs are replaced with lib9p client 3. Spawn/ProcessKeeper spawn 9P server
This commit is contained in:
parent
6fc7cf999b
commit
427cadeb29
|
|
@ -1,43 +1,48 @@
|
|||
// toolsrv — tool execution server.
|
||||
// toolsrv — 9P tool execution server.
|
||||
//
|
||||
// Accepts JSON-RPC 2.0 connections. Executes tool calls in a sandboxed
|
||||
// environment. Used both locally (spawned by olliesrv) and remotely
|
||||
// (deployed via SSH bootstrap).
|
||||
// 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]
|
||||
// toolsrv serve --cwd /path/to/project --listen /path/to/sock [--yolo]
|
||||
//
|
||||
// When --listen is provided, accepts connections on a Unix socket.
|
||||
// Otherwise, serves a single session over stdin/stdout (bootstrap mode).
|
||||
// Listens on a Unix socket and serves a 9P filesystem for tool execution.
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"flag"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"syscall"
|
||||
|
||||
"ollie/env"
|
||||
"ollie/fsedsl"
|
||||
"ollie/toolsrv"
|
||||
)
|
||||
|
||||
var (
|
||||
cwd = flag.String("cwd", ".", "working directory for execution")
|
||||
listenPath = flag.String("listen", "", "Unix socket path to listen on (default: stdio mode)")
|
||||
listenPath = flag.String("listen", "", "Unix socket path to listen on (required)")
|
||||
yolo = flag.Bool("yolo", false, "skip sandbox enforcement")
|
||||
)
|
||||
|
||||
func main() {
|
||||
if len(os.Args) < 2 || os.Args[1] != "serve" {
|
||||
fmt.Fprintln(os.Stderr, "usage: toolsrv serve --cwd <path> [--listen <socket>]")
|
||||
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)
|
||||
}
|
||||
|
||||
env.EnsureDefaults()
|
||||
|
||||
// Prepend our bin dir to PATH so landrun is found (deployed alongside us).
|
||||
|
|
@ -45,35 +50,34 @@ func main() {
|
|||
binDir := filepath.Join(home, ".config", "ollie", "bin")
|
||||
os.Setenv("PATH", binDir+":"+os.Getenv("PATH"))
|
||||
|
||||
// Get shared secret from environment
|
||||
secret := os.Getenv("TOOLSRV_SECRET")
|
||||
if secret == "" {
|
||||
fmt.Fprintln(os.Stderr, "error: TOOLSRV_SECRET environment variable required")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// Create tool registry
|
||||
toolReg, _ := toolsrv.NewRegistry()
|
||||
sessionID := os.Getenv("OLLIE_SESSION_ID")
|
||||
|
||||
var opts []toolsrv.Option
|
||||
if *yolo {
|
||||
opts = append(opts, toolsrv.WithYolo())
|
||||
}
|
||||
// Create 9P server
|
||||
srv := toolsrv.NewServer9P([]byte(secret))
|
||||
srv.SetYolo(*yolo)
|
||||
if toolReg != nil && sessionID != "" {
|
||||
opts = append(opts, toolsrv.WithToolRegistry(toolReg, sessionID))
|
||||
}
|
||||
server := toolsrv.New(*cwd)
|
||||
for _, o := range opts {
|
||||
o(server)
|
||||
srv.SetRegistry(toolReg, sessionID)
|
||||
}
|
||||
|
||||
// Wire up environment propagation for late-arriving session IDs.
|
||||
server.OnEnvSet = func(key, value string) {
|
||||
switch key {
|
||||
case "OLLIE_SESSION_ID":
|
||||
os.Setenv("OLLIE_SESSION_ID", value)
|
||||
if toolReg != nil {
|
||||
server.SetToolRegistry(toolReg, value)
|
||||
}
|
||||
case "OLLIE_UNAME":
|
||||
os.Setenv("OLLIE_UNAME", value)
|
||||
}
|
||||
// Build the filesystem tree
|
||||
ctx := toolsrv.ToolsrvCtx{Server: srv}
|
||||
tree := fsedsl.BuildTree(toolsrv.ToolsrvSpec(), ctx)
|
||||
if tree == nil {
|
||||
fmt.Fprintln(os.Stderr, "error: failed to build filesystem tree")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// 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)
|
||||
|
|
@ -82,12 +86,48 @@ func main() {
|
|||
cancel()
|
||||
}()
|
||||
|
||||
if *listenPath != "" {
|
||||
if err := server.ServeSocket(ctx, *listenPath); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "%v\n", err)
|
||||
os.Exit(1)
|
||||
// Remove stale socket if it exists
|
||||
os.Remove(*listenPath)
|
||||
|
||||
// Listen on Unix socket
|
||||
ln, err := net.Listen("unix", *listenPath)
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "error: listen: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer ln.Close()
|
||||
defer os.Remove(*listenPath)
|
||||
|
||||
fmt.Fprintf(os.Stderr, "toolsrv: listening on %s\n", *listenPath)
|
||||
|
||||
// Accept connections
|
||||
go func() {
|
||||
<-runCtx.Done()
|
||||
ln.Close()
|
||||
}()
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
select {
|
||||
case <-runCtx.Done():
|
||||
return
|
||||
default:
|
||||
fmt.Fprintf(os.Stderr, "accept error: %v\n", err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
} else {
|
||||
server.ServeRPC(ctx, os.Stdin, os.Stdout)
|
||||
go serve9P(runCtx, conn, tree)
|
||||
}
|
||||
}
|
||||
|
||||
// serve9P handles a single 9P connection.
|
||||
// TODO: This needs to implement the 9P protocol using the fsedsl tree.
|
||||
// For now this is a placeholder - the actual implementation should follow
|
||||
// the pattern in cmd/olliesrv/server.go.
|
||||
func serve9P(ctx context.Context, conn net.Conn, tree *fsedsl.Tree) {
|
||||
defer conn.Close()
|
||||
// TODO: Implement 9P protocol handling
|
||||
// This would use plan9.ReadFcall/WriteFcall and dispatch to tree operations
|
||||
fmt.Fprintf(os.Stderr, "toolsrv: new connection (9P serving not yet implemented)\n")
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,212 @@
|
|||
// compat.go - Compatibility types during 9P migration.
|
||||
// These stub out the old JSON-RPC interface while we transition to 9P.
|
||||
// TODO: Remove once 9P client is fully implemented.
|
||||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// Process represents a spawned toolsrv subprocess.
|
||||
// TODO: Update for 9P model.
|
||||
type Process struct {
|
||||
Socket string
|
||||
Info HostInfo
|
||||
cleanup func()
|
||||
}
|
||||
|
||||
// Close shuts down the process.
|
||||
func (p *Process) Close() {
|
||||
if p.cleanup != nil {
|
||||
p.cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
// ProcessKeeper manages toolsrv process lifecycle and reconnection.
|
||||
// TODO: Update for 9P model.
|
||||
type ProcessKeeper struct {
|
||||
proc *Process
|
||||
conn *Conn
|
||||
ctx context.Context
|
||||
respawn func(ctx context.Context) (*Process, error)
|
||||
}
|
||||
|
||||
// NewProcessKeeper creates a new ProcessKeeper.
|
||||
func NewProcessKeeper(ctx context.Context, proc *Process, respawn func(ctx context.Context) (*Process, error)) *ProcessKeeper {
|
||||
return &ProcessKeeper{proc: proc, ctx: ctx, respawn: respawn}
|
||||
}
|
||||
|
||||
// Dial returns a connection to the toolsrv.
|
||||
func (pk *ProcessKeeper) Dial() (*Conn, error) {
|
||||
// TODO: Implement actual 9P connection
|
||||
if pk.conn == nil {
|
||||
pk.conn = &Conn{}
|
||||
}
|
||||
return pk.conn, nil
|
||||
}
|
||||
|
||||
// SetContext updates the context.
|
||||
func (pk *ProcessKeeper) SetContext(ctx context.Context) {
|
||||
pk.ctx = ctx
|
||||
}
|
||||
|
||||
// Close shuts down the keeper.
|
||||
func (pk *ProcessKeeper) Close() {
|
||||
// TODO: Implement
|
||||
}
|
||||
|
||||
// Spawn starts a local toolsrv subprocess.
|
||||
// TODO: Update to spawn 9P server.
|
||||
func Spawn(ctx context.Context, cwd string, opts ...Option) (*Process, error) {
|
||||
// TODO: Implement
|
||||
return &Process{}, nil
|
||||
}
|
||||
|
||||
// SpawnRemote starts a toolsrv on a remote host via SSH.
|
||||
// TODO: Update to spawn 9P server.
|
||||
func SpawnRemote(ctx context.Context, cfg RemoteConfig) (*Process, error) {
|
||||
// TODO: Implement
|
||||
return &Process{}, nil
|
||||
}
|
||||
|
||||
// RemoteConfig holds configuration for remote toolsrv.
|
||||
type RemoteConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
User string
|
||||
CWD string
|
||||
SSHTarget string
|
||||
}
|
||||
|
||||
// Option is a functional option for toolsrv.
|
||||
type Option func(*optionState)
|
||||
|
||||
type optionState struct {
|
||||
yolo bool
|
||||
}
|
||||
|
||||
// WithYolo returns an option to skip sandbox enforcement.
|
||||
func WithYolo() Option {
|
||||
return func(o *optionState) {
|
||||
o.yolo = true
|
||||
}
|
||||
}
|
||||
|
||||
// HostInfo contains information about the host running toolsrv.
|
||||
type HostInfo struct {
|
||||
Platform string
|
||||
IsGitRepo bool
|
||||
}
|
||||
|
||||
// FetchHostInfo retrieves host information from toolsrv.
|
||||
func FetchHostInfo(conn *Conn) (HostInfo, error) {
|
||||
// TODO: Implement via 9P read to /info
|
||||
return HostInfo{}, nil
|
||||
}
|
||||
|
||||
// Conn is a stub for the old RPC connection.
|
||||
// TODO: Replace with 9P client using lib9p.
|
||||
type Conn struct {
|
||||
// Placeholder - will be replaced with 9P fsys
|
||||
}
|
||||
|
||||
// Runner interface for tool execution.
|
||||
type Runner interface {
|
||||
ListTools() ([]ToolInfo, error)
|
||||
CallTool(ctx context.Context, name string, args json.RawMessage) (json.RawMessage, error)
|
||||
}
|
||||
|
||||
// ListTools returns the list of available tools.
|
||||
func (c *Conn) ListTools() ([]ToolInfo, error) {
|
||||
// TODO: Implement via 9P read to /tools
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// CallTool calls a tool and returns the result.
|
||||
func (c *Conn) CallTool(ctx context.Context, name string, args json.RawMessage) (json.RawMessage, error) {
|
||||
// TODO: Implement via 9P rdwr to /proc/new
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// SetEnv sets an environment variable on the remote server.
|
||||
func (c *Conn) SetEnv(key, value string) {
|
||||
// TODO: Implement via 9P (or remove - may not be needed with 9P model)
|
||||
}
|
||||
|
||||
// SetCWD sets the working directory on the remote server.
|
||||
func (c *Conn) SetCWD(dir string) {
|
||||
// TODO: Implement via 9P (or remove - may not be needed with 9P model)
|
||||
}
|
||||
|
||||
// SetOnToolsChanged sets a callback for when tools change.
|
||||
func (c *Conn) SetOnToolsChanged(fn func()) {
|
||||
// TODO: Implement or remove
|
||||
}
|
||||
|
||||
// ToolRegistryRevision returns the current revision of the tool registry.
|
||||
func (c *Conn) ToolRegistryRevision() uint64 {
|
||||
// TODO: Implement or remove
|
||||
return 0
|
||||
}
|
||||
|
||||
// Ping checks if the connection is alive.
|
||||
func (c *Conn) Ping() error {
|
||||
// TODO: Implement via 9P
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the connection.
|
||||
func (c *Conn) Close() {
|
||||
// TODO: Implement
|
||||
}
|
||||
|
||||
// Detach detaches a running process.
|
||||
func (c *Conn) Detach() bool {
|
||||
// TODO: Implement via 9P
|
||||
return false
|
||||
}
|
||||
|
||||
// ListDetachedRaw returns raw info about detached processes.
|
||||
func (c *Conn) ListDetachedRaw() []any {
|
||||
// TODO: Implement via 9P read to /proc
|
||||
return nil
|
||||
}
|
||||
|
||||
// SignalDetached sends a signal to a detached process.
|
||||
func (c *Conn) SignalDetached(pid int, sig syscall.Signal) error {
|
||||
// TODO: Implement via 9P write to /proc/{pid}/ctl
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetDetachedOutput gets output from a detached process.
|
||||
func (c *Conn) GetDetachedOutput(pid int) (string, error) {
|
||||
// TODO: Implement via 9P read to /proc/{pid}/out
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// DismissDetached dismisses a detached process.
|
||||
func (c *Conn) DismissDetached(pid int) bool {
|
||||
// TODO: Implement via 9P write to /proc/{pid}/ctl
|
||||
return false
|
||||
}
|
||||
|
||||
// WithOutputStream returns a context with streaming output function attached.
|
||||
func WithOutputStream(ctx context.Context, fn func(string)) context.Context {
|
||||
return WithStreamFunc(ctx, fn)
|
||||
}
|
||||
|
||||
// RateLimitedError indicates a rate limit was hit.
|
||||
type RateLimitedError struct {
|
||||
Err error
|
||||
Remaining int
|
||||
}
|
||||
|
||||
func (e RateLimitedError) Error() string {
|
||||
return e.Err.Error()
|
||||
}
|
||||
|
||||
// Add MediaType and Data fields to ToolResultContent if needed
|
||||
// by extending the existing type (already defined in exec9p.go).
|
||||
// Note: These may need to be in the actual ToolResultContent struct.
|
||||
304
toolsrv/conn.go
304
toolsrv/conn.go
|
|
@ -1,304 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// Conn is a unified JSON-RPC client connected to a tool server subprocess
|
||||
// (ollie-remote). It satisfies Runner and all extended interfaces that the
|
||||
// agent layer uses via type assertions. Both local and remote connections
|
||||
// return the same *Conn — the only difference is the underlying transport.
|
||||
type Conn struct {
|
||||
writeMu sync.Mutex
|
||||
pendingMu sync.Mutex
|
||||
pending map[int64]*pendingCall
|
||||
closed chan struct{}
|
||||
closeOnce sync.Once
|
||||
readErr error
|
||||
|
||||
enc *json.Encoder
|
||||
dec *json.Decoder
|
||||
nextID atomic.Int64
|
||||
Info HostInfo
|
||||
|
||||
// cleanup is called on Close (kills subprocess, closes SSH, etc.)
|
||||
cleanup func()
|
||||
|
||||
// onToolsChanged is a callback for tool directory change notifications.
|
||||
onToolsChanged func()
|
||||
}
|
||||
|
||||
// NewConn wraps an io.ReadWriteCloser as a JSON-RPC connection.
|
||||
// The cleanup function is called when Close() is invoked.
|
||||
func NewConn(rwc io.ReadWriteCloser, cleanup func()) *Conn {
|
||||
c := &Conn{
|
||||
enc: json.NewEncoder(rwc),
|
||||
dec: json.NewDecoder(bufio.NewReader(rwc)),
|
||||
cleanup: cleanup,
|
||||
pending: make(map[int64]*pendingCall),
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
go c.readLoop()
|
||||
return c
|
||||
}
|
||||
|
||||
// NewConnSplit creates a Conn from separate reader/writer (e.g. stdin/stdout pipes).
|
||||
func NewConnSplit(r io.Reader, w io.WriteCloser, cleanup func()) *Conn {
|
||||
c := &Conn{
|
||||
enc: json.NewEncoder(w),
|
||||
dec: json.NewDecoder(bufio.NewReader(r)),
|
||||
cleanup: cleanup,
|
||||
pending: make(map[int64]*pendingCall),
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
go c.readLoop()
|
||||
return c
|
||||
}
|
||||
|
||||
// --- Runner interface ---
|
||||
|
||||
func (c *Conn) ListTools() ([]ToolInfo, error) {
|
||||
resp, err := c.call("list_tools", nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var infos []ToolInfo
|
||||
if err := json.Unmarshal(resp, &infos); err != nil {
|
||||
return nil, fmt.Errorf("list_tools unmarshal: %w", err)
|
||||
}
|
||||
return infos, nil
|
||||
}
|
||||
|
||||
func (c *Conn) CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) {
|
||||
return c.callStreaming(ctx, tool, args)
|
||||
}
|
||||
|
||||
// --- Extended interfaces (used by agent via type assertion) ---
|
||||
|
||||
func (c *Conn) SetEnv(key, value string) {
|
||||
params, _ := json.Marshal(map[string]string{"key": key, "value": value})
|
||||
c.call("set_env", params)
|
||||
}
|
||||
|
||||
func (c *Conn) SetCWD(dir string) {
|
||||
params, _ := json.Marshal(map[string]string{"dir": dir})
|
||||
c.call("set_cwd", params)
|
||||
}
|
||||
|
||||
func (c *Conn) SetOnToolsChanged(fn func()) {
|
||||
c.pendingMu.Lock()
|
||||
c.onToolsChanged = fn
|
||||
c.pendingMu.Unlock()
|
||||
}
|
||||
|
||||
// ToolRegistryRevision returns 0 for remote connections (revision tracking
|
||||
// is not supported over RPC; use SetOnToolsChanged for push-based refresh).
|
||||
func (c *Conn) ToolRegistryRevision() uint64 { return 0 }
|
||||
|
||||
func (c *Conn) Detach() bool {
|
||||
resp, err := c.call("detach", nil)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
var ok bool
|
||||
json.Unmarshal(resp, &ok)
|
||||
return ok
|
||||
}
|
||||
|
||||
func (c *Conn) ListDetachedRaw() []any {
|
||||
resp, err := c.call("list_detached", nil)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
var out []any
|
||||
json.Unmarshal(resp, &out)
|
||||
return out
|
||||
}
|
||||
|
||||
func (c *Conn) SignalDetached(pid int, sig syscall.Signal) error {
|
||||
params, _ := json.Marshal(map[string]any{"pid": pid, "signal": int(sig)})
|
||||
_, err := c.call("signal_detached", params)
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *Conn) GetDetachedOutput(pid int) (string, error) {
|
||||
params, _ := json.Marshal(map[string]any{"pid": pid})
|
||||
resp, err := c.call("get_detached_output", params)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
var out string
|
||||
json.Unmarshal(resp, &out)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *Conn) DismissDetached(pid int) bool {
|
||||
params, _ := json.Marshal(map[string]any{"pid": pid})
|
||||
resp, err := c.call("dismiss_detached", params)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
var ok bool
|
||||
json.Unmarshal(resp, &ok)
|
||||
return ok
|
||||
}
|
||||
|
||||
// Close shuts down the connection and subprocess.
|
||||
func (c *Conn) Close() {
|
||||
c.fail(fmt.Errorf("connection closed"))
|
||||
if c.cleanup != nil {
|
||||
c.cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
// Ping verifies the remote server is responsive.
|
||||
func (c *Conn) Ping() error {
|
||||
_, err := c.call("ping", nil)
|
||||
return err
|
||||
}
|
||||
|
||||
// FetchHostInfo retrieves environment details from the server.
|
||||
func (c *Conn) FetchHostInfo() (HostInfo, error) {
|
||||
resp, err := c.call("host_info", nil)
|
||||
if err != nil {
|
||||
return HostInfo{}, err
|
||||
}
|
||||
var info HostInfo
|
||||
if err := json.Unmarshal(resp, &info); err != nil {
|
||||
return HostInfo{}, err
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
type pendingCall struct {
|
||||
response chan rpcResponse
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func (c *Conn) readLoop() {
|
||||
for {
|
||||
var resp rpcResponse
|
||||
if err := c.dec.Decode(&resp); err != nil {
|
||||
c.fail(fmt.Errorf("rpc read: %w", err))
|
||||
return
|
||||
}
|
||||
|
||||
c.pendingMu.Lock()
|
||||
call := c.pending[resp.ID]
|
||||
c.pendingMu.Unlock()
|
||||
if call == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Tool output is sent as an intermediate response with the request ID.
|
||||
var notif outputNotification
|
||||
if resp.Stream && resp.Error == nil && json.Unmarshal(resp.Result, ¬if) == nil {
|
||||
if notif.Data != "" {
|
||||
StreamOutput(call.ctx, notif.Data)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
c.pendingMu.Lock()
|
||||
delete(c.pending, resp.ID)
|
||||
c.pendingMu.Unlock()
|
||||
call.response <- resp
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) fail(err error) {
|
||||
c.closeOnce.Do(func() {
|
||||
c.pendingMu.Lock()
|
||||
c.readErr = err
|
||||
close(c.closed)
|
||||
for id, call := range c.pending {
|
||||
delete(c.pending, id)
|
||||
call.response <- rpcResponse{Error: &rpcError{Code: -32000, Message: err.Error()}}
|
||||
}
|
||||
c.pendingMu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
// --- JSON-RPC internals ---
|
||||
|
||||
func (c *Conn) call(method string, params json.RawMessage) (json.RawMessage, error) {
|
||||
return c.callStreaming(context.Background(), method, params)
|
||||
}
|
||||
|
||||
func (c *Conn) callStreaming(ctx context.Context, method string, params json.RawMessage) (json.RawMessage, error) {
|
||||
id := c.nextID.Add(1)
|
||||
call := &pendingCall{response: make(chan rpcResponse, 1), ctx: ctx}
|
||||
c.pendingMu.Lock()
|
||||
select {
|
||||
case <-c.closed:
|
||||
c.pendingMu.Unlock()
|
||||
return nil, fmt.Errorf("rpc %s: connection closed", method)
|
||||
default:
|
||||
}
|
||||
c.pending[id] = call
|
||||
c.pendingMu.Unlock()
|
||||
|
||||
c.writeMu.Lock()
|
||||
err := c.enc.Encode(rpcRequest{JSONRPC: "2.0", ID: id, Method: method, Params: params})
|
||||
c.writeMu.Unlock()
|
||||
if err != nil {
|
||||
c.pendingMu.Lock()
|
||||
delete(c.pending, id)
|
||||
c.pendingMu.Unlock()
|
||||
return nil, fmt.Errorf("rpc write %s: %w", method, err)
|
||||
}
|
||||
|
||||
select {
|
||||
case resp := <-call.response:
|
||||
if resp.Error != nil {
|
||||
return nil, fmt.Errorf("rpc %s: %s", method, resp.Error.Message)
|
||||
}
|
||||
return resp.Result, nil
|
||||
case <-ctx.Done():
|
||||
c.pendingMu.Lock()
|
||||
delete(c.pending, id)
|
||||
c.pendingMu.Unlock()
|
||||
return nil, ctx.Err()
|
||||
case <-c.closed:
|
||||
c.pendingMu.Lock()
|
||||
err := c.readErr
|
||||
c.pendingMu.Unlock()
|
||||
if err == nil {
|
||||
err = fmt.Errorf("connection closed")
|
||||
}
|
||||
return nil, fmt.Errorf("rpc %s: %w", method, err)
|
||||
}
|
||||
}
|
||||
|
||||
// JSON-RPC 2.0 types shared between Conn (client) and transport layer.
|
||||
|
||||
type rpcRequest struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int64 `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
type rpcResponse struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int64 `json:"id"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
Error *rpcError `json:"error,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
}
|
||||
|
||||
type rpcError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
type outputNotification struct {
|
||||
Data string `json:"data"`
|
||||
}
|
||||
|
|
@ -1,99 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const RingBufSize = 64 * 1024 // 64KB ring buffer per detached process
|
||||
|
||||
// DetachedProcess represents a process that the agent has detached from
|
||||
// but which continues running. The user can view its output and signal it.
|
||||
type DetachedProcess struct {
|
||||
PID int
|
||||
Command string
|
||||
Started time.Time
|
||||
Exited bool
|
||||
ExitCode int
|
||||
|
||||
Ring *RingBuffer
|
||||
Done chan struct{}
|
||||
Mu sync.Mutex
|
||||
}
|
||||
|
||||
// Info returns a plain-data snapshot of this process for external consumers.
|
||||
type InfoData struct {
|
||||
PID int
|
||||
Command string
|
||||
Started int64 // unix timestamp
|
||||
Exited bool
|
||||
ExitCode int
|
||||
}
|
||||
|
||||
func (p *DetachedProcess) Info() InfoData {
|
||||
p.Mu.Lock()
|
||||
defer p.Mu.Unlock()
|
||||
return InfoData{
|
||||
PID: p.PID,
|
||||
Command: p.Command,
|
||||
Started: p.Started.Unix(),
|
||||
Exited: p.Exited,
|
||||
ExitCode: p.ExitCode,
|
||||
}
|
||||
}
|
||||
|
||||
// Output returns the current contents of the ring buffer.
|
||||
func (p *DetachedProcess) Output() string {
|
||||
p.Mu.Lock()
|
||||
defer p.Mu.Unlock()
|
||||
return p.Ring.String()
|
||||
}
|
||||
|
||||
// ringBuffer is a fixed-size circular byte buffer.
|
||||
type RingBuffer struct {
|
||||
buf []byte
|
||||
size int
|
||||
pos int
|
||||
full bool
|
||||
}
|
||||
|
||||
func NewRingBuffer(size int) *RingBuffer {
|
||||
return &RingBuffer{buf: make([]byte, size), size: size}
|
||||
}
|
||||
|
||||
// Write implements io.Writer.
|
||||
func (r *RingBuffer) Write(p []byte) (int, error) {
|
||||
n := len(p)
|
||||
if n >= r.size {
|
||||
// Data larger than buffer: just keep the tail
|
||||
copy(r.buf, p[n-r.size:])
|
||||
r.pos = 0
|
||||
r.full = true
|
||||
return n, nil
|
||||
}
|
||||
if r.pos+n <= r.size {
|
||||
copy(r.buf[r.pos:], p)
|
||||
} else {
|
||||
first := r.size - r.pos
|
||||
copy(r.buf[r.pos:], p[:first])
|
||||
copy(r.buf, p[first:])
|
||||
r.full = true
|
||||
}
|
||||
r.pos = (r.pos + n) % r.size
|
||||
if r.pos == 0 && n > 0 {
|
||||
r.full = true
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// String returns the buffer contents in order.
|
||||
func (r *RingBuffer) String() string {
|
||||
if !r.full {
|
||||
return string(r.buf[:r.pos])
|
||||
}
|
||||
// Buffer has wrapped: data from pos..end + 0..pos
|
||||
out := make([]byte, r.size)
|
||||
copy(out, r.buf[r.pos:])
|
||||
copy(out[r.size-r.pos:], r.buf[:r.pos])
|
||||
return string(out)
|
||||
}
|
||||
|
|
@ -1,25 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
)
|
||||
|
||||
// Dial connects to a running tool server at the given Unix socket path.
|
||||
// The server must already be running (started via Spawn or SpawnRemote).
|
||||
func Dial(ctx context.Context, socket string) (*Conn, error) {
|
||||
if socket == "" {
|
||||
return nil, fmt.Errorf("toolsrv.Dial: socket path required")
|
||||
}
|
||||
conn, err := net.Dial("unix", socket)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("toolsrv.Dial: %w", err)
|
||||
}
|
||||
c := NewConn(conn, func() { conn.Close() })
|
||||
if err := c.Ping(); err != nil {
|
||||
c.Close()
|
||||
return nil, fmt.Errorf("toolsrv.Dial: ping failed: %w", err)
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
594
toolsrv/exec.go
594
toolsrv/exec.go
|
|
@ -1,594 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"ollie/paths"
|
||||
"ollie/sandbox"
|
||||
)
|
||||
|
||||
func loadSandboxConfig(name string) (*sandbox.Config, error) {
|
||||
if name == "" {
|
||||
name = "default"
|
||||
}
|
||||
path := filepath.Join(paths.CfgDir(), "sandbox", name+".yaml")
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("sandbox %q not found: %w", name, err)
|
||||
}
|
||||
defer f.Close()
|
||||
return sandbox.LoadSandbox(f)
|
||||
}
|
||||
|
||||
type limitedWriter struct {
|
||||
mu sync.Mutex
|
||||
w io.Writer
|
||||
written int
|
||||
limit int
|
||||
truncated bool
|
||||
stream func(string) // if non-nil, called with each chunk of output
|
||||
}
|
||||
|
||||
func (lw *limitedWriter) Write(p []byte) (n int, err error) {
|
||||
lw.mu.Lock()
|
||||
defer lw.mu.Unlock()
|
||||
if lw.written >= lw.limit {
|
||||
lw.truncated = true
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
remaining := lw.limit - lw.written
|
||||
toWrite := p
|
||||
if len(p) > remaining {
|
||||
toWrite = p[:remaining]
|
||||
lw.truncated = true
|
||||
}
|
||||
|
||||
written, err := lw.w.Write(toWrite)
|
||||
lw.written += written
|
||||
if lw.stream != nil {
|
||||
lw.stream(string(toWrite[:written]))
|
||||
}
|
||||
if err != nil {
|
||||
return written, err
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// Connects to the broker socket, sends the request with the current env,
|
||||
// and streams the framed response back.
|
||||
func (s *Server) executeBypass(ctx context.Context, cmd, dir string, timeout int, doDetach ...bool) (string, error) {
|
||||
return s.executeBypassOpts(ctx, cmd, dir, timeout, false, doDetach...)
|
||||
}
|
||||
|
||||
func (s *Server) executeBypassSudo(ctx context.Context, cmd, dir string, timeout int) (string, error) {
|
||||
return s.executeBypassOpts(ctx, cmd, dir, timeout, true)
|
||||
}
|
||||
|
||||
func (s *Server) executeBypassOpts(ctx context.Context, cmd, dir string, timeout int, sudo bool, doDetach ...bool) (string, error) {
|
||||
xdg := os.Getenv("XDG_RUNTIME_DIR")
|
||||
if xdg == "" {
|
||||
return "", fmt.Errorf("bypass not available: no XDG_RUNTIME_DIR")
|
||||
}
|
||||
sockPath := filepath.Join(xdg, "ollie", "bypass.sock")
|
||||
|
||||
wantDetach := len(doDetach) > 0 && doDetach[0]
|
||||
|
||||
var cancel context.CancelFunc
|
||||
if wantDetach || timeout <= 0 {
|
||||
ctx, cancel = context.WithCancel(ctx)
|
||||
} else {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
||||
}
|
||||
|
||||
// Connect to broker
|
||||
conn, err := net.DialTimeout("unix", sockPath, 5*time.Second)
|
||||
if err != nil {
|
||||
cancel()
|
||||
return "", fmt.Errorf("bypass not available: %w", err)
|
||||
}
|
||||
|
||||
// Send request — merge process env with session-scoped envExtra.
|
||||
envMap := make(map[string]string, len(os.Environ()))
|
||||
for _, kv := range os.Environ() {
|
||||
if k, v, ok := strings.Cut(kv, "="); ok {
|
||||
envMap[k] = v
|
||||
}
|
||||
}
|
||||
s.envMu.RLock()
|
||||
for k, v := range s.envExtra {
|
||||
envMap[k] = v
|
||||
}
|
||||
s.envMu.RUnlock()
|
||||
reqJSON, _ := json.Marshal(struct {
|
||||
Cmd string `json:"cmd"`
|
||||
Cwd string `json:"cwd"`
|
||||
Env map[string]string `json:"env"`
|
||||
Session string `json:"session,omitempty"`
|
||||
Sudo bool `json:"sudo,omitempty"`
|
||||
}{Cmd: cmd, Cwd: dir, Env: envMap, Session: s.sessionID, Sudo: sudo})
|
||||
reqJSON = append(reqJSON, '\n')
|
||||
if _, err := conn.Write(reqJSON); err != nil {
|
||||
conn.Close()
|
||||
cancel()
|
||||
return "", fmt.Errorf("bypass execution failed: write: %w", err)
|
||||
}
|
||||
|
||||
// Set up detach channel
|
||||
detachCh := make(chan struct{}, 1)
|
||||
s.detachMu.Lock()
|
||||
s.detachCh = detachCh
|
||||
s.detachMu.Unlock()
|
||||
defer func() {
|
||||
s.detachMu.Lock()
|
||||
s.detachCh = nil
|
||||
s.detachMu.Unlock()
|
||||
}()
|
||||
|
||||
if wantDetach {
|
||||
close(detachCh)
|
||||
}
|
||||
|
||||
// readFrames reads from the connection, writing to w and streaming.
|
||||
// Returns exit code when 'x' frame arrives or -1 on error.
|
||||
readFrames := func(w io.Writer, stream func(string)) int {
|
||||
header := make([]byte, 5)
|
||||
for {
|
||||
conn.SetReadDeadline(time.Now().Add(1 * time.Second)) //nolint:errcheck
|
||||
_, err := io.ReadFull(conn, header)
|
||||
if err != nil {
|
||||
if ne, ok := err.(net.Error); ok && ne.Timeout() {
|
||||
// Check for context cancellation
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return -1
|
||||
default:
|
||||
continue
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
frameType := header[0]
|
||||
length := binary.BigEndian.Uint32(header[1:5])
|
||||
|
||||
payload := make([]byte, length)
|
||||
if length > 0 {
|
||||
conn.SetReadDeadline(time.Now().Add(30 * time.Second)) //nolint:errcheck
|
||||
if _, err := io.ReadFull(conn, payload); err != nil {
|
||||
return -1
|
||||
}
|
||||
}
|
||||
|
||||
switch frameType {
|
||||
case 'd':
|
||||
w.Write(payload) //nolint:errcheck
|
||||
if stream != nil {
|
||||
stream(string(payload))
|
||||
}
|
||||
case 'x':
|
||||
var code int
|
||||
fmt.Sscanf(string(payload), "%d", &code)
|
||||
return code
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check if detach was requested
|
||||
select {
|
||||
case <-detachCh:
|
||||
// Detach: background the socket read into a goroutine
|
||||
cmdStr := "bypass: " + cmd
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := NewRingBuffer(RingBufSize)
|
||||
pid := int(time.Now().UnixNano() & 0x7FFFFFFF) // synthetic PID
|
||||
|
||||
proc := &DetachedProcess{
|
||||
PID: pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
s.detachMu.Lock()
|
||||
s.detached = append(s.detached, proc)
|
||||
s.detachMu.Unlock()
|
||||
|
||||
go func() {
|
||||
defer conn.Close()
|
||||
defer cancel()
|
||||
exitCode := readFrames(ring, nil)
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
proc.ExitCode = exitCode
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if s.OnExit != nil {
|
||||
s.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if s.OnDetach != nil {
|
||||
s.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
return fmt.Sprintf("[detached: pid %d]", pid), nil
|
||||
|
||||
default:
|
||||
// Normal (foreground) execution — must also handle manual detach.
|
||||
var outputBuf bytes.Buffer
|
||||
streamFn := StreamFunc(ctx)
|
||||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
limit: 10 * 1024 * 1024,
|
||||
stream: streamFn,
|
||||
}
|
||||
|
||||
// Run readFrames in a goroutine so we can select on detach.
|
||||
type frameResult struct{ exitCode int }
|
||||
frameCh := make(chan frameResult, 1)
|
||||
go func() {
|
||||
code := readFrames(lw, nil)
|
||||
frameCh <- frameResult{code}
|
||||
}()
|
||||
|
||||
select {
|
||||
case fr := <-frameCh:
|
||||
// Normal completion
|
||||
conn.Close()
|
||||
cancel()
|
||||
combined := outputBuf.String()
|
||||
if fr.exitCode != 0 {
|
||||
if combined == "" {
|
||||
return "", fmt.Errorf("bypass execution failed (exit %d)", fr.exitCode)
|
||||
}
|
||||
errOutput := combined
|
||||
const maxErrOutput = 8192
|
||||
if len(errOutput) > maxErrOutput {
|
||||
errOutput = errOutput[:maxErrOutput] + fmt.Sprintf("\n[...truncated %d bytes in error]", len(combined)-maxErrOutput)
|
||||
}
|
||||
return combined, fmt.Errorf("bypass execution failed (exit %d)\nOutput: %s", fr.exitCode, errOutput)
|
||||
}
|
||||
return combined, nil
|
||||
|
||||
case <-ctx.Done():
|
||||
// Timeout or external cancellation
|
||||
conn.Close()
|
||||
cancel()
|
||||
combined := outputBuf.String()
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return combined, fmt.Errorf("bypass execution timeout after %d seconds", timeout)
|
||||
}
|
||||
return combined, fmt.Errorf("bypass execution interrupted")
|
||||
|
||||
case <-detachCh:
|
||||
// Manual detach: background the connection read
|
||||
cmdStr := "bypass: " + cmd
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := NewRingBuffer(RingBufSize)
|
||||
pid := int(time.Now().UnixNano() & 0x7FFFFFFF)
|
||||
|
||||
// Splice: future output goes to ring buffer, stop streaming
|
||||
lw.mu.Lock()
|
||||
lw.w = ring
|
||||
lw.stream = nil
|
||||
lw.mu.Unlock()
|
||||
|
||||
proc := &DetachedProcess{
|
||||
PID: pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
s.detachMu.Lock()
|
||||
s.detached = append(s.detached, proc)
|
||||
s.detachMu.Unlock()
|
||||
|
||||
go func() {
|
||||
fr := <-frameCh
|
||||
conn.Close()
|
||||
cancel()
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
proc.ExitCode = fr.exitCode
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if s.OnExit != nil {
|
||||
s.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if s.OnDetach != nil {
|
||||
s.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
|
||||
partial := outputBuf.String()
|
||||
return partial + fmt.Sprintf("\n[detached: pid %d]", pid), nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// executeWithStdin runs code in a sandbox and returns combined stdout+stderr.
|
||||
// For languages where code is itself passed via stdin (ed, expect, bc), stdinData is ignored.
|
||||
func (s *Server) executeWithStdin(ctx context.Context, code, language string, timeout int, sandboxName string, trusted bool, stdinData string, doDetach ...bool) (string, error) {
|
||||
if timeout < 0 {
|
||||
timeout = 30
|
||||
}
|
||||
|
||||
if !trusted {
|
||||
if err := s.ValidateCode(code, language); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
var cfg *sandbox.Config
|
||||
var err error
|
||||
if !s.Yolo {
|
||||
cfg, err = loadSandboxConfig(sandboxName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
s.wdMu.RLock()
|
||||
workDir := s.cwd
|
||||
s.wdMu.RUnlock()
|
||||
if workDir == "" {
|
||||
workDir, _ = os.Getwd()
|
||||
}
|
||||
|
||||
var cancel context.CancelFunc
|
||||
if timeout > 0 {
|
||||
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
|
||||
} else {
|
||||
// timeout=0 means no timeout (run indefinitely, e.g. for daemons)
|
||||
ctx, cancel = context.WithCancel(ctx)
|
||||
}
|
||||
// NOTE: do NOT defer cancel() here — the detach path must avoid cancelling
|
||||
// the context (which would kill the detached process via cmd.Cancel).
|
||||
// Each exit path calls cancel() explicitly except detach.
|
||||
|
||||
var cmd *exec.Cmd
|
||||
var interpreter []string
|
||||
// codeStdin: non-empty means the code itself is fed via stdin (ed, expect, bc).
|
||||
// In these cases stdinData cannot be used simultaneously.
|
||||
var codeStdin string
|
||||
switch language {
|
||||
case "bash", "":
|
||||
interpreter = []string{"bash", "-c", code}
|
||||
default:
|
||||
cancel()
|
||||
return "", fmt.Errorf("unsupported language: %s (only bash is supported)", language)
|
||||
}
|
||||
s.envMu.RLock()
|
||||
envMap := make(map[string]string, len(s.envExtra))
|
||||
for k, v := range s.envExtra {
|
||||
envMap[k] = v
|
||||
}
|
||||
s.envMu.RUnlock()
|
||||
for _, ev := range os.Environ() {
|
||||
k, v, _ := strings.Cut(ev, "=")
|
||||
if _, ok := envMap[k]; !ok {
|
||||
envMap[k] = v
|
||||
}
|
||||
}
|
||||
if _, ok := envMap["NAMESPACE"]; !ok {
|
||||
if ns := plan9Namespace(envMap); ns != "" {
|
||||
envMap["NAMESPACE"] = ns
|
||||
}
|
||||
}
|
||||
getenv := func(key string) string { return envMap[key] }
|
||||
|
||||
if s.Yolo {
|
||||
cmd = exec.CommandContext(ctx, interpreter[0], interpreter[1:]...)
|
||||
} else {
|
||||
wrapped, wrapErr := sandbox.WrapCommand(cfg, interpreter, workDir, getenv)
|
||||
if wrapErr != nil {
|
||||
cancel()
|
||||
return "", wrapErr
|
||||
}
|
||||
cmd = exec.CommandContext(ctx, wrapped[0], wrapped[1:]...)
|
||||
}
|
||||
cmd.Dir = workDir
|
||||
switch {
|
||||
case codeStdin != "":
|
||||
cmd.Stdin = strings.NewReader(codeStdin)
|
||||
case stdinData != "":
|
||||
cmd.Stdin = strings.NewReader(stdinData)
|
||||
}
|
||||
|
||||
s.envMu.RLock()
|
||||
baseEnv := os.Environ()
|
||||
// Remove keys that envExtra overrides so duplicates don't leak through.
|
||||
filtered := baseEnv[:0]
|
||||
for _, ev := range baseEnv {
|
||||
k, _, _ := strings.Cut(ev, "=")
|
||||
if _, overridden := s.envExtra[k]; !overridden {
|
||||
filtered = append(filtered, ev)
|
||||
}
|
||||
}
|
||||
cmd.Env = filtered
|
||||
for k, v := range s.envExtra {
|
||||
cmd.Env = append(cmd.Env, k+"="+v)
|
||||
}
|
||||
s.envMu.RUnlock()
|
||||
// Export computed NAMESPACE so tools like ollie-9p can find the server.
|
||||
if ns := envMap["NAMESPACE"]; ns != "" {
|
||||
cmd.Env = append(cmd.Env, "NAMESPACE="+ns)
|
||||
}
|
||||
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
cmd.Cancel = func() error {
|
||||
if cmd.Process != nil {
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
cmd.WaitDelay = time.Second
|
||||
|
||||
var outputBuf bytes.Buffer
|
||||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
limit: 10 * 1024 * 1024,
|
||||
stream: StreamFunc(ctx),
|
||||
}
|
||||
cmd.Stdout = lw
|
||||
cmd.Stderr = lw
|
||||
|
||||
// Set up detach channel for this execution
|
||||
detachCh := make(chan struct{}, 1)
|
||||
s.detachMu.Lock()
|
||||
s.detachCh = detachCh
|
||||
s.detachMu.Unlock()
|
||||
defer func() {
|
||||
s.detachMu.Lock()
|
||||
s.detachCh = nil
|
||||
s.detachMu.Unlock()
|
||||
}()
|
||||
|
||||
if err = cmd.Start(); err != nil {
|
||||
cancel()
|
||||
return "", fmt.Errorf("execution failed: %w", err)
|
||||
}
|
||||
|
||||
// If detach requested, immediately signal the detach channel
|
||||
if len(doDetach) > 0 && doDetach[0] {
|
||||
close(detachCh)
|
||||
}
|
||||
|
||||
// Wait for completion, context cancellation, or detach signal
|
||||
waitCh := make(chan error, 1)
|
||||
go func() { waitCh <- cmd.Wait() }()
|
||||
|
||||
select {
|
||||
case err = <-waitCh:
|
||||
// Normal completion
|
||||
cancel()
|
||||
output := outputBuf.Bytes()
|
||||
if lw.truncated {
|
||||
output = append(output, []byte("\n[output truncated at 10MB]")...)
|
||||
}
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return "", fmt.Errorf("execution timeout after %d seconds", timeout)
|
||||
}
|
||||
if err != nil {
|
||||
// Cap output in error message to avoid flooding the agent context.
|
||||
// The full output is still returned as the first return value.
|
||||
errOutput := string(output)
|
||||
const maxErrOutput = 8192
|
||||
if len(errOutput) > maxErrOutput {
|
||||
errOutput = errOutput[:maxErrOutput] + fmt.Sprintf("\n[...truncated %d bytes in error]", len(output)-maxErrOutput)
|
||||
}
|
||||
return string(output), fmt.Errorf("execution failed: %w\nOutput: %s", err, errOutput)
|
||||
}
|
||||
return string(output), nil
|
||||
|
||||
case <-ctx.Done():
|
||||
// Context cancelled (interrupt or timeout)
|
||||
cancel()
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
<-waitCh // reap
|
||||
output := outputBuf.Bytes()
|
||||
if ctx.Err() == context.DeadlineExceeded {
|
||||
return string(output), fmt.Errorf("execution timeout after %d seconds", timeout)
|
||||
}
|
||||
return string(output), fmt.Errorf("execution interrupted")
|
||||
|
||||
case <-detachCh:
|
||||
// Detach: move process to registry, continue running
|
||||
// Neutralize cmd.Cancel and WaitDelay so context cancellation
|
||||
// won't kill the process after the timeout fires.
|
||||
cmd.Cancel = nil
|
||||
cmd.WaitDelay = 0
|
||||
cancel()
|
||||
|
||||
cmdStr := strings.Join(interpreter, " ")
|
||||
if len(cmdStr) > 80 {
|
||||
cmdStr = cmdStr[:77] + "..."
|
||||
}
|
||||
ring := NewRingBuffer(RingBufSize)
|
||||
// Splice: future output goes to ring buffer instead of outputBuf.
|
||||
lw.mu.Lock()
|
||||
lw.w = ring
|
||||
lw.stream = nil
|
||||
lw.mu.Unlock()
|
||||
|
||||
proc := &DetachedProcess{
|
||||
PID: cmd.Process.Pid,
|
||||
Command: cmdStr,
|
||||
Started: time.Now(),
|
||||
Ring: ring,
|
||||
Done: make(chan struct{}),
|
||||
}
|
||||
s.detachMu.Lock()
|
||||
s.detached = append(s.detached, proc)
|
||||
s.detachMu.Unlock()
|
||||
|
||||
// Monitor for exit in background
|
||||
go func() {
|
||||
waitErr := <-waitCh
|
||||
proc.Mu.Lock()
|
||||
proc.Exited = true
|
||||
if waitErr != nil {
|
||||
if exitErr, ok := waitErr.(*exec.ExitError); ok {
|
||||
proc.ExitCode = exitErr.ExitCode()
|
||||
} else {
|
||||
proc.ExitCode = -1
|
||||
}
|
||||
}
|
||||
proc.Mu.Unlock()
|
||||
close(proc.Done)
|
||||
if s.OnExit != nil {
|
||||
s.OnExit(proc.PID, proc.ExitCode)
|
||||
}
|
||||
}()
|
||||
|
||||
if s.OnDetach != nil {
|
||||
s.OnDetach(proc.PID, cmdStr)
|
||||
}
|
||||
|
||||
partial := outputBuf.String()
|
||||
return partial + fmt.Sprintf("\n[detached: pid %d]", proc.PID), nil
|
||||
}
|
||||
}
|
||||
|
||||
// plan9Namespace computes the plan9port namespace directory from the environment.
|
||||
// Convention: $NAMESPACE if set, else /tmp/ns.<user>.<display> where display is
|
||||
// $DISPLAY with trailing .0 stripped and / replaced by _.
|
||||
func plan9Namespace(env map[string]string) string {
|
||||
if ns := env["NAMESPACE"]; ns != "" {
|
||||
return ns
|
||||
}
|
||||
user := env["USER"]
|
||||
if user == "" {
|
||||
return ""
|
||||
}
|
||||
disp := env["DISPLAY"]
|
||||
if disp == "" {
|
||||
return ""
|
||||
}
|
||||
// Canonicalize: strip trailing .0
|
||||
if strings.HasSuffix(disp, ".0") {
|
||||
disp = disp[:len(disp)-2]
|
||||
}
|
||||
// Replace / with _
|
||||
disp = strings.ReplaceAll(disp, "/", "_")
|
||||
return "/tmp/ns." + user + "." + disp
|
||||
}
|
||||
|
|
@ -29,6 +29,35 @@ type ExecConfig struct {
|
|||
Timeout int
|
||||
}
|
||||
|
||||
// ToolResult is the JSON structure returned by tool execution.
|
||||
type ToolResult struct {
|
||||
Content []ToolResultContent `json:"content"`
|
||||
IsError bool `json:"isError,omitempty"`
|
||||
}
|
||||
|
||||
// ToolResultContent is a single content item in a tool result.
|
||||
type ToolResultContent struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text,omitempty"`
|
||||
MediaType string `json:"mediaType,omitempty"`
|
||||
Data string `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
// StreamFunc returns the streaming output function from context, if any.
|
||||
func StreamFunc(ctx context.Context) func(string) {
|
||||
if fn, ok := ctx.Value(streamKey{}).(func(string)); ok {
|
||||
return fn
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type streamKey struct{}
|
||||
|
||||
// WithStreamFunc attaches a streaming output function to context.
|
||||
func WithStreamFunc(ctx context.Context, fn func(string)) context.Context {
|
||||
return context.WithValue(ctx, streamKey{}, fn)
|
||||
}
|
||||
|
||||
// ExecuteTool runs a tool script with the given args and returns the result.
|
||||
// This is the core execution function, decoupled from Server.
|
||||
func ExecuteTool(ctx context.Context, info ToolInfo, args json.RawMessage, cfg ExecConfig) (json.RawMessage, error) {
|
||||
|
|
@ -316,3 +345,49 @@ func readBypassFrames(ctx context.Context, conn net.Conn, w *bytes.Buffer, strea
|
|||
}
|
||||
}
|
||||
}
|
||||
|
||||
// plan9Namespace computes the Plan 9 namespace directory.
|
||||
func plan9Namespace(env map[string]string) string {
|
||||
if ns := env["NAMESPACE"]; ns != "" {
|
||||
return ns
|
||||
}
|
||||
disp := env["DISPLAY"]
|
||||
if disp == "" {
|
||||
return "/tmp/ns." + env["USER"] + "." + "unix"
|
||||
}
|
||||
// Strip leading "localhost" if present
|
||||
disp = strings.TrimPrefix(disp, "localhost")
|
||||
disp = strings.ReplaceAll(disp, "/", "_")
|
||||
return "/tmp/ns." + env["USER"] + "." + disp
|
||||
}
|
||||
|
||||
// limitedWriter wraps a writer with a size limit and optional streaming.
|
||||
type limitedWriter struct {
|
||||
w io.Writer
|
||||
limit int
|
||||
written int
|
||||
truncated bool
|
||||
stream func(string)
|
||||
}
|
||||
|
||||
func (lw *limitedWriter) Write(p []byte) (int, error) {
|
||||
if lw.truncated {
|
||||
return len(p), nil // discard but report success
|
||||
}
|
||||
remaining := lw.limit - lw.written
|
||||
if remaining <= 0 {
|
||||
lw.truncated = true
|
||||
return len(p), nil
|
||||
}
|
||||
toWrite := p
|
||||
if len(p) > remaining {
|
||||
toWrite = p[:remaining]
|
||||
lw.truncated = true
|
||||
}
|
||||
n, err := lw.w.Write(toWrite)
|
||||
lw.written += n
|
||||
if lw.stream != nil && n > 0 {
|
||||
lw.stream(string(toWrite[:n]))
|
||||
}
|
||||
return len(p), err
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,112 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ProcessKeeper maintains a running ollie-remote process and respawns it
|
||||
// on failure. Multiple agents can Dial() the same keeper to get independent
|
||||
// connections. If the process dies, the next Dial() triggers a respawn.
|
||||
type ProcessKeeper struct {
|
||||
mu sync.Mutex
|
||||
ctx context.Context
|
||||
proc *Process
|
||||
spawn func(context.Context) (*Process, error)
|
||||
}
|
||||
|
||||
// NewProcessKeeper creates a keeper with an initial process and a spawn factory.
|
||||
func NewProcessKeeper(ctx context.Context, proc *Process, spawn func(context.Context) (*Process, error)) *ProcessKeeper {
|
||||
return &ProcessKeeper{
|
||||
ctx: ctx,
|
||||
proc: proc,
|
||||
spawn: spawn,
|
||||
}
|
||||
}
|
||||
|
||||
// Dial returns a new Conn to the managed process. If the process is dead
|
||||
// (socket unreachable), it respawns before dialing.
|
||||
func (pk *ProcessKeeper) Dial() (*Conn, error) {
|
||||
pk.mu.Lock()
|
||||
sock := pk.proc.Socket
|
||||
pk.mu.Unlock()
|
||||
|
||||
// Try dialing the existing socket.
|
||||
conn, err := Dial(pk.ctx, sock)
|
||||
if err == nil {
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// Dial failed — process likely dead. Respawn.
|
||||
return pk.respawnAndDial()
|
||||
}
|
||||
|
||||
// Socket returns the current process socket path.
|
||||
func (pk *ProcessKeeper) Socket() string {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
return pk.proc.Socket
|
||||
}
|
||||
|
||||
// Process returns the current underlying Process.
|
||||
func (pk *ProcessKeeper) Process() *Process {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
return pk.proc
|
||||
}
|
||||
|
||||
// Close shuts down the managed process.
|
||||
func (pk *ProcessKeeper) Close() {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
if pk.proc != nil {
|
||||
pk.proc.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// SetContext updates the keeper's context. Used on session resume to
|
||||
// provide a fresh context after the previous one was cancelled on pause.
|
||||
func (pk *ProcessKeeper) SetContext(ctx context.Context) {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
pk.ctx = ctx
|
||||
}
|
||||
|
||||
func (pk *ProcessKeeper) respawnAndDial() (*Conn, error) {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
|
||||
// Another goroutine may have respawned already — try current socket first.
|
||||
if conn, err := Dial(pk.ctx, pk.proc.Socket); err == nil {
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// Clean up old process.
|
||||
pk.proc.Close()
|
||||
|
||||
var lastErr error
|
||||
for attempt := 0; attempt < 3; attempt++ {
|
||||
if pk.ctx.Err() != nil {
|
||||
return nil, pk.ctx.Err()
|
||||
}
|
||||
proc, err := pk.spawn(pk.ctx)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
time.Sleep(time.Duration(attempt+1) * 2 * time.Second)
|
||||
continue
|
||||
}
|
||||
pk.proc = proc
|
||||
|
||||
conn, err := Dial(pk.ctx, proc.Socket)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
proc.Close()
|
||||
time.Sleep(time.Duration(attempt+1) * 2 * time.Second)
|
||||
continue
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
return nil, fmt.Errorf("respawn failed after 3 attempts: %w", lastErr)
|
||||
}
|
||||
|
|
@ -1,52 +0,0 @@
|
|||
// Package toolsrv remote types.
|
||||
//
|
||||
// RemoteConfig and HostInfo are used by the transport layer and session
|
||||
// creation. The actual client implementation lives in conn.go; the
|
||||
// transport (SSH bootstrap) lives in transport.go.
|
||||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
)
|
||||
|
||||
// HostInfo holds environment details from the remote host.
|
||||
type HostInfo struct {
|
||||
Platform string `json:"platform"`
|
||||
Arch string `json:"arch"`
|
||||
IsGitRepo bool `json:"is_git_repo"`
|
||||
}
|
||||
|
||||
// RemoteConfig holds the parameters for connecting to a remote host.
|
||||
type RemoteConfig struct {
|
||||
// SSHTarget is the SSH destination (e.g., "user@host" or an SSH config alias).
|
||||
SSHTarget string
|
||||
// CWD is the working directory on the remote host.
|
||||
CWD string
|
||||
// ToolsPath overrides the remote tool scripts path.
|
||||
ToolsPath string
|
||||
// Yolo disables sandbox enforcement on the remote.
|
||||
Yolo bool
|
||||
}
|
||||
|
||||
// RemoteDial opens an SSH connection to the remote host and returns a *Conn.
|
||||
// Convenience wrapper: spawns the remote process and dials a single connection.
|
||||
func RemoteDial(ctx context.Context, cfg RemoteConfig) (*Conn, error) {
|
||||
proc, err := SpawnRemote(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := Dial(ctx, proc.Socket)
|
||||
if err != nil {
|
||||
proc.Close()
|
||||
return nil, err
|
||||
}
|
||||
origClose := conn.cleanup
|
||||
conn.cleanup = func() {
|
||||
if origClose != nil {
|
||||
origClose()
|
||||
}
|
||||
proc.Close()
|
||||
}
|
||||
conn.Info = proc.Info
|
||||
return conn, nil
|
||||
}
|
||||
|
|
@ -1,322 +0,0 @@
|
|||
package toolsrv_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
"syscall"
|
||||
"testing"
|
||||
|
||||
"ollie/toolsrv"
|
||||
)
|
||||
|
||||
// --- Test infrastructure ---
|
||||
|
||||
// rpcRequest mirrors the JSON-RPC 2.0 request format (test-local).
|
||||
type rpcRequest struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int64 `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
// rpcResponse mirrors the JSON-RPC 2.0 response format (test-local).
|
||||
type rpcResponse struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int64 `json:"id"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
Error *rpcError `json:"error,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
}
|
||||
|
||||
type rpcError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// testServer wraps a toolsrv.Server with a pipe-based RPC connection.
|
||||
type testServer struct {
|
||||
srv *toolsrv.Server
|
||||
conn net.Conn // client side
|
||||
enc *json.Encoder
|
||||
dec *json.Decoder
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func newTestServer(t *testing.T) *testServer {
|
||||
t.Helper()
|
||||
cwd := t.TempDir()
|
||||
os.MkdirAll(filepath.Join(cwd, ".git"), 0755)
|
||||
|
||||
reg, err := toolsrv.NewRegistry()
|
||||
if err != nil {
|
||||
t.Fatalf("NewRegistry: %v", err)
|
||||
}
|
||||
|
||||
sessionID := "test-session-" + t.Name()
|
||||
srv := toolsrv.New(cwd)
|
||||
srv.SetToolRegistry(reg, sessionID)
|
||||
|
||||
clientConn, serverConn := net.Pipe()
|
||||
ts := &testServer{
|
||||
srv: srv,
|
||||
conn: clientConn,
|
||||
enc: json.NewEncoder(clientConn),
|
||||
dec: json.NewDecoder(clientConn),
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
ts.wg.Add(1)
|
||||
go func() {
|
||||
defer ts.wg.Done()
|
||||
srv.ServeRPC(ctx, serverConn, serverConn)
|
||||
}()
|
||||
|
||||
t.Cleanup(func() {
|
||||
cancel()
|
||||
clientConn.Close()
|
||||
serverConn.Close()
|
||||
ts.wg.Wait()
|
||||
})
|
||||
|
||||
return ts
|
||||
}
|
||||
|
||||
// call sends an RPC request and returns the response (skipping stream frames).
|
||||
func (ts *testServer) call(method string, params any) rpcResponse {
|
||||
var raw json.RawMessage
|
||||
if params != nil {
|
||||
raw, _ = json.Marshal(params)
|
||||
}
|
||||
ts.enc.Encode(rpcRequest{JSONRPC: "2.0", ID: 1, Method: method, Params: raw})
|
||||
for {
|
||||
var resp rpcResponse
|
||||
ts.dec.Decode(&resp)
|
||||
if !resp.Stream {
|
||||
return resp
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Integration tests ---
|
||||
|
||||
func TestRPC_Ping(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("ping", nil)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
var result string
|
||||
json.Unmarshal(resp.Result, &result)
|
||||
if result != "pong" {
|
||||
t.Errorf("expected 'pong', got %q", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_HostInfo(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("host_info", nil)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
|
||||
var info toolsrv.HostInfo
|
||||
json.Unmarshal(resp.Result, &info)
|
||||
if info.Platform != runtime.GOOS {
|
||||
t.Errorf("platform: expected %q, got %q", runtime.GOOS, info.Platform)
|
||||
}
|
||||
if info.Arch != runtime.GOARCH {
|
||||
t.Errorf("arch: expected %q, got %q", runtime.GOARCH, info.Arch)
|
||||
}
|
||||
if !info.IsGitRepo {
|
||||
t.Error("expected is_git_repo=true (we created .git dir)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_ListTools_Empty(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("list_tools", nil)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
|
||||
var tools []toolsrv.ToolInfo
|
||||
json.Unmarshal(resp.Result, &tools)
|
||||
if len(tools) != 0 {
|
||||
t.Errorf("expected 0 tools, got %d", len(tools))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_SetEnv(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
|
||||
var envKey, envValue string
|
||||
ts.srv.OnEnvSet = func(k, v string) {
|
||||
envKey, envValue = k, v
|
||||
}
|
||||
|
||||
resp := ts.call("set_env", map[string]string{"key": "TEST_VAR", "value": "test_value"})
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
if envKey != "TEST_VAR" || envValue != "test_value" {
|
||||
t.Errorf("env not set: got key=%q value=%q", envKey, envValue)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_SetCWD(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("set_cwd", map[string]string{"dir": "/tmp/test-cwd"})
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
var result bool
|
||||
json.Unmarshal(resp.Result, &result)
|
||||
if !result {
|
||||
t.Error("expected result=true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_Detach_NoRunningProcess(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("detach", nil)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
var result bool
|
||||
json.Unmarshal(resp.Result, &result)
|
||||
if result {
|
||||
t.Error("expected false when no process running")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_ListDetached_Empty(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("list_detached", nil)
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
var list []any
|
||||
json.Unmarshal(resp.Result, &list)
|
||||
if len(list) != 0 {
|
||||
t.Errorf("expected empty list, got %d items", len(list))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_SignalDetached_NotFound(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("signal_detached", map[string]any{"pid": 99999, "signal": int(syscall.SIGTERM)})
|
||||
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected error for nonexistent PID")
|
||||
}
|
||||
if resp.Error.Code != -32000 {
|
||||
t.Errorf("expected error code -32000, got %d", resp.Error.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_GetDetachedOutput_NotFound(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("get_detached_output", map[string]any{"pid": 99999})
|
||||
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected error for nonexistent PID")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_DismissDetached_NotFound(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("dismiss_detached", map[string]any{"pid": 99999})
|
||||
|
||||
if resp.Error != nil {
|
||||
t.Fatalf("unexpected error: %v", resp.Error.Message)
|
||||
}
|
||||
var result bool
|
||||
json.Unmarshal(resp.Result, &result)
|
||||
if result {
|
||||
t.Error("expected false for nonexistent PID")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_CallTool_NotLoaded(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
resp := ts.call("nonexistent_tool", map[string]any{"arg": "value"})
|
||||
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected error for unloaded tool")
|
||||
}
|
||||
if resp.Error.Code != -32000 {
|
||||
t.Errorf("expected error code -32000, got %d", resp.Error.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRPC_InvalidParams(t *testing.T) {
|
||||
ts := newTestServer(t)
|
||||
// Send an array where an object is expected — json.Unmarshal will fail.
|
||||
ts.enc.Encode(rpcRequest{JSONRPC: "2.0", ID: 1, Method: "set_env", Params: json.RawMessage(`[1,2,3]`)})
|
||||
for {
|
||||
var resp rpcResponse
|
||||
if err := ts.dec.Decode(&resp); err != nil {
|
||||
t.Fatalf("decode: %v", err)
|
||||
}
|
||||
if resp.Stream {
|
||||
continue
|
||||
}
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected error for invalid params")
|
||||
}
|
||||
if resp.Error.Code != -32602 {
|
||||
t.Errorf("expected error code -32602 (invalid params), got %d", resp.Error.Code)
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// --- Full pipe-based client/server tests ---
|
||||
|
||||
func TestRPC_PipeConnection(t *testing.T) {
|
||||
cwd := t.TempDir()
|
||||
os.MkdirAll(filepath.Join(cwd, ".git"), 0755)
|
||||
|
||||
reg, _ := toolsrv.NewRegistry()
|
||||
srv := toolsrv.New(cwd)
|
||||
srv.SetToolRegistry(reg, "pipe-test")
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer clientConn.Close()
|
||||
defer serverConn.Close()
|
||||
|
||||
go srv.ServeRPC(ctx, serverConn, serverConn)
|
||||
|
||||
// Use toolsrv.Conn as the client
|
||||
conn := toolsrv.NewConn(clientConn, nil)
|
||||
|
||||
if err := conn.Ping(); err != nil {
|
||||
t.Errorf("Ping failed: %v", err)
|
||||
}
|
||||
|
||||
tools, err := conn.ListTools()
|
||||
if err != nil {
|
||||
t.Errorf("ListTools failed: %v", err)
|
||||
}
|
||||
if len(tools) != 0 {
|
||||
t.Errorf("expected 0 tools, got %d", len(tools))
|
||||
}
|
||||
|
||||
conn.SetEnv("TEST_KEY", "TEST_VALUE")
|
||||
conn.SetCWD("/tmp/test")
|
||||
}
|
||||
|
|
@ -1,263 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"ollie/paths"
|
||||
)
|
||||
|
||||
// RPCResponse is a JSON-RPC 2.0 response.
|
||||
type RPCResponse struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id"`
|
||||
Result any `json:"result,omitempty"`
|
||||
Error *RPCError `json:"error,omitempty"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
}
|
||||
|
||||
// lockedEncoder serializes JSON writes to a single writer.
|
||||
type lockedEncoder struct {
|
||||
mu sync.Mutex
|
||||
enc *json.Encoder
|
||||
}
|
||||
|
||||
func (e *lockedEncoder) Encode(v any) error {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
return e.enc.Encode(v)
|
||||
}
|
||||
|
||||
// ServeSocket accepts JSON-RPC connections on a Unix socket. Blocks until
|
||||
// ctx is cancelled.
|
||||
func (s *Server) ServeSocket(ctx context.Context, sockPath string) error {
|
||||
os.Remove(sockPath)
|
||||
os.MkdirAll(filepath.Dir(sockPath), 0700)
|
||||
|
||||
ln, err := net.Listen("unix", sockPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("listen %s: %w", sockPath, err)
|
||||
}
|
||||
defer ln.Close()
|
||||
defer os.Remove(sockPath)
|
||||
|
||||
fmt.Fprintln(os.Stderr, "ListenReady")
|
||||
|
||||
var mu sync.Mutex
|
||||
var activeConns []net.Conn
|
||||
var wg sync.WaitGroup
|
||||
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
ln.Close()
|
||||
mu.Lock()
|
||||
for _, c := range activeConns {
|
||||
c.Close()
|
||||
}
|
||||
mu.Unlock()
|
||||
}()
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
if ctx.Err() != nil {
|
||||
break
|
||||
}
|
||||
continue
|
||||
}
|
||||
mu.Lock()
|
||||
activeConns = append(activeConns, conn)
|
||||
mu.Unlock()
|
||||
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
defer conn.Close()
|
||||
s.ServeRPC(ctx, conn, conn)
|
||||
mu.Lock()
|
||||
for i, c := range activeConns {
|
||||
if c == conn {
|
||||
activeConns = append(activeConns[:i], activeConns[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
mu.Unlock()
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
return nil
|
||||
}
|
||||
|
||||
// ServeRPC reads JSON-RPC 2.0 requests from in and writes responses to out.
|
||||
// Blocks until in is closed or ctx is cancelled.
|
||||
func (s *Server) ServeRPC(ctx context.Context, in io.Reader, out io.Writer) {
|
||||
dec := json.NewDecoder(in)
|
||||
enc := &lockedEncoder{enc: json.NewEncoder(out)}
|
||||
|
||||
for {
|
||||
var req RPCRequest
|
||||
if err := dec.Decode(&req); err != nil {
|
||||
if err == io.EOF || ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
enc.Encode(RPCResponse{
|
||||
JSONRPC: "2.0",
|
||||
ID: req.ID,
|
||||
Error: &RPCError{Code: -32700, Message: "parse error: " + err.Error()},
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
go s.handleRPC(ctx, req, enc)
|
||||
}
|
||||
}
|
||||
|
||||
// CWD returns the server's current working directory.
|
||||
func (s *Server) CWD() string {
|
||||
s.wdMu.RLock()
|
||||
defer s.wdMu.RUnlock()
|
||||
return s.cwd
|
||||
}
|
||||
|
||||
// handleRPC dispatches a single JSON-RPC request.
|
||||
//
|
||||
// Reserved method names (control plane):
|
||||
// - list_tools, host_info, ping, set_env, set_cwd
|
||||
// - detach, list_detached, signal_detached, get_detached_output, dismiss_detached
|
||||
// - tool_load
|
||||
//
|
||||
// Any other method name is treated as a tool invocation (data plane).
|
||||
// A tool with the same name as a reserved method is unreachable.
|
||||
func (s *Server) handleRPC(ctx context.Context, req RPCRequest, enc *lockedEncoder) {
|
||||
switch req.Method {
|
||||
case "list_tools":
|
||||
tools, err := s.ListTools()
|
||||
if err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32000, Message: err.Error()}})
|
||||
} else {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: tools})
|
||||
}
|
||||
|
||||
case "host_info":
|
||||
info := HostInfo{
|
||||
Platform: runtime.GOOS,
|
||||
Arch: runtime.GOARCH,
|
||||
IsGitRepo: paths.IsGitRepo(s.CWD()),
|
||||
}
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: info})
|
||||
|
||||
case "ping":
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: "pong"})
|
||||
|
||||
case "set_env":
|
||||
var params struct {
|
||||
Key string `json:"key"`
|
||||
Value string `json:"value"`
|
||||
}
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32602, Message: err.Error()}})
|
||||
return
|
||||
}
|
||||
s.SetEnv(params.Key, params.Value)
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: true})
|
||||
|
||||
case "set_cwd":
|
||||
var params struct {
|
||||
Dir string `json:"dir"`
|
||||
}
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32602, Message: err.Error()}})
|
||||
return
|
||||
}
|
||||
s.SetCWD(params.Dir)
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: true})
|
||||
|
||||
case "detach":
|
||||
ok := s.Detach()
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: ok})
|
||||
|
||||
case "list_detached":
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: s.ListDetachedRaw()})
|
||||
|
||||
case "signal_detached":
|
||||
var params struct {
|
||||
PID int `json:"pid"`
|
||||
Signal int `json:"signal"`
|
||||
}
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32602, Message: err.Error()}})
|
||||
return
|
||||
}
|
||||
if err := s.SignalDetached(params.PID, syscall.Signal(params.Signal)); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32000, Message: err.Error()}})
|
||||
} else {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: true})
|
||||
}
|
||||
|
||||
case "tool_load":
|
||||
var params struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32602, Message: err.Error()}})
|
||||
return
|
||||
}
|
||||
if err := s.LoadTool(params.Name); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32000, Message: err.Error()}})
|
||||
} else {
|
||||
res, _ := json.Marshal(map[string]string{"result": "loaded: " + params.Name})
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: json.RawMessage(res)})
|
||||
}
|
||||
|
||||
case "get_detached_output":
|
||||
var params struct {
|
||||
PID int `json:"pid"`
|
||||
}
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32602, Message: err.Error()}})
|
||||
return
|
||||
}
|
||||
out, err := s.GetDetachedOutput(params.PID)
|
||||
if err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32000, Message: err.Error()}})
|
||||
} else {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: out})
|
||||
}
|
||||
|
||||
case "dismiss_detached":
|
||||
var params struct {
|
||||
PID int `json:"pid"`
|
||||
}
|
||||
if err := json.Unmarshal(req.Params, ¶ms); err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32602, Message: err.Error()}})
|
||||
return
|
||||
}
|
||||
ok := s.DismissDetached(params.PID)
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: ok})
|
||||
|
||||
default:
|
||||
// Data plane: method name is the tool name.
|
||||
streamCtx := WithOutputStream(ctx, func(data string) {
|
||||
enc.Encode(RPCResponse{
|
||||
JSONRPC: "2.0",
|
||||
ID: req.ID,
|
||||
Result: OutputNotification{Data: data},
|
||||
Stream: true,
|
||||
})
|
||||
})
|
||||
result, err := s.CallTool(streamCtx, req.Method, req.Params)
|
||||
if err != nil {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Error: &RPCError{Code: -32000, Message: err.Error()}})
|
||||
} else {
|
||||
enc.Encode(RPCResponse{JSONRPC: "2.0", ID: req.ID, Result: json.RawMessage(result)})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,51 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import "encoding/json"
|
||||
|
||||
// RPCRequest is a JSON-RPC 2.0 request.
|
||||
type RPCRequest struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
// RPCError is a JSON-RPC 2.0 error object.
|
||||
type RPCError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// OutputNotification carries streamed tool output data.
|
||||
type OutputNotification struct {
|
||||
Data string `json:"data"`
|
||||
}
|
||||
|
||||
// ToolResult is the JSON envelope returned by tool execution.
|
||||
// Tool scripts write this to stdout; the agent loop parses it via
|
||||
// extractToolResult in agent/runtime.go.
|
||||
//
|
||||
// Schema:
|
||||
//
|
||||
// {
|
||||
// "isError": false,
|
||||
// "content": [
|
||||
// {"type": "text", "text": "..."},
|
||||
// {"type": "image", "media_type": "image/png", "data": "<base64>"}
|
||||
// ]
|
||||
// }
|
||||
//
|
||||
// If a tool emits raw text (not JSON), the agent treats it as a single
|
||||
// text content block with isError=false.
|
||||
type ToolResult struct {
|
||||
IsError bool `json:"isError,omitempty"`
|
||||
Content []ToolResultContent `json:"content"`
|
||||
}
|
||||
|
||||
// ToolResultContent is a single content block within a ToolResult.
|
||||
type ToolResultContent struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text,omitempty"`
|
||||
MediaType string `json:"media_type,omitempty"`
|
||||
Data string `json:"data,omitempty"`
|
||||
}
|
||||
|
|
@ -1,369 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"ollie/paths"
|
||||
)
|
||||
|
||||
const (
|
||||
failureWindow = 1 * time.Minute
|
||||
maxFailures = 5
|
||||
blockDuration = 5 * time.Minute
|
||||
)
|
||||
|
||||
// RateLimitedError is returned by checkRateLimit when the shell is blocked
|
||||
// due to too many validation failures. The agent loop uses this to detect
|
||||
// the condition and inject a clear signal rather than counting it as a
|
||||
// normal tool error.
|
||||
type RateLimitedError struct {
|
||||
Remaining time.Duration
|
||||
}
|
||||
|
||||
func (e *RateLimitedError) Error() string {
|
||||
return fmt.Sprintf("rate limited: too many validation failures, blocked for %v", e.Remaining)
|
||||
}
|
||||
|
||||
// Server runs code in a sandboxed environment.
|
||||
type Server struct {
|
||||
// cwd is the working directory for sandboxed commands. If empty,
|
||||
// the process working directory is used.
|
||||
wdMu sync.RWMutex
|
||||
cwd string
|
||||
|
||||
// envExtra holds per-session environment variables injected via SetEnv.
|
||||
envMu sync.RWMutex
|
||||
envExtra map[string]string
|
||||
|
||||
// Hooks for lifecycle events (harnesses like 9P can inject mount logic here)
|
||||
OnClose func()
|
||||
OnEnvSet func(key, value string)
|
||||
|
||||
// Yolo skips the landrun sandbox for all execution.
|
||||
Yolo bool
|
||||
|
||||
toolRegistry *Registry
|
||||
sessionID string
|
||||
|
||||
// OnInjection is called when a skill is loaded and its content
|
||||
// should be injected into the agent's context.
|
||||
OnInjection func(content string)
|
||||
|
||||
// OnToolsChanged is called when the set of loaded tools changes.
|
||||
// The receiver should re-fetch the tool list to update its state.
|
||||
OnToolsChanged func()
|
||||
|
||||
// rate limiting state (per-Server)
|
||||
rateLimitMu sync.Mutex
|
||||
validationFailures int
|
||||
lastFailure time.Time
|
||||
blockedUntil time.Time
|
||||
|
||||
// Detached process management
|
||||
detachMu sync.Mutex
|
||||
detachCh chan struct{} // signal to detach the currently running process
|
||||
detached []*DetachedProcess
|
||||
OnDetach func(pid int, cmd string) // hook: called when a process is detached
|
||||
OnExit func(pid int, exitCode int) // hook: called when a detached process exits
|
||||
}
|
||||
|
||||
// Option configures a Server.
|
||||
type Option func(*Server)
|
||||
|
||||
// WithYolo skips the landrun sandbox.
|
||||
func WithYolo() Option { return func(s *Server) { s.Yolo = true } }
|
||||
|
||||
|
||||
|
||||
// WithToolRegistry attaches a tool registry and session ID to the Server.
|
||||
func WithToolRegistry(r *Registry, sessionID string) Option {
|
||||
return func(s *Server) {
|
||||
s.toolRegistry = r
|
||||
s.sessionID = sessionID
|
||||
}
|
||||
}
|
||||
|
||||
// LoadTool loads a tool by name into this server's tool registry.
|
||||
// Returns an error if no registry is configured or the tool cannot be found.
|
||||
// If OnToolsChanged is set, it is called to signal that the available tools
|
||||
// have changed. The caller is responsible for re-fetching the current list.
|
||||
func (s *Server) LoadTool(name string) error {
|
||||
if s.toolRegistry == nil {
|
||||
return fmt.Errorf("no tool registry configured")
|
||||
}
|
||||
if s.sessionID == "" {
|
||||
return fmt.Errorf("no session ID configured")
|
||||
}
|
||||
if err := s.toolRegistry.Load(s.sessionID, name); err != nil {
|
||||
return err
|
||||
}
|
||||
if s.OnToolsChanged != nil {
|
||||
s.OnToolsChanged()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
|
||||
// ListTools implements Server, returning tools loaded in the session's
|
||||
// tool registry. With zero built-in tools, only dynamically loaded tools
|
||||
// appear here.
|
||||
func (s *Server) ListTools() ([]ToolInfo, error) {
|
||||
var all []ToolInfo
|
||||
|
||||
if s.toolRegistry != nil && s.sessionID != "" {
|
||||
all = append(all, s.toolRegistry.Loaded(s.sessionID)...)
|
||||
}
|
||||
return all, nil
|
||||
}
|
||||
|
||||
// CallTool implements Runner.
|
||||
// With zero built-in tools, all tools must be loaded in the registry first.
|
||||
func (s *Server) CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) {
|
||||
if s.toolRegistry != nil && s.sessionID != "" {
|
||||
if _, promoted := s.toolRegistry.Lookup(s.sessionID, tool); promoted {
|
||||
return s.callPromotedTool(ctx, tool, args)
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("unknown tool: %s (not loaded)", tool)
|
||||
}
|
||||
|
||||
// New creates a new Server with the given working directory.
|
||||
func New(cwd string) *Server { return &Server{cwd: paths.ExpandHome(cwd)} }
|
||||
|
||||
// SetCWD updates the working directory used for subsequent command executions.
|
||||
func (s *Server) SetCWD(dir string) {
|
||||
s.wdMu.Lock()
|
||||
s.cwd = paths.ExpandHome(dir)
|
||||
s.wdMu.Unlock()
|
||||
}
|
||||
|
||||
// SetEnv adds a session-scoped environment variable injected into all
|
||||
// subsequent subprocess invocations for this session.
|
||||
func (s *Server) SetEnv(key, value string) {
|
||||
s.envMu.Lock()
|
||||
if s.envExtra == nil {
|
||||
s.envExtra = make(map[string]string)
|
||||
}
|
||||
s.envExtra[key] = value
|
||||
s.envMu.Unlock()
|
||||
if s.OnEnvSet != nil {
|
||||
s.OnEnvSet(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
// SetToolRegistry attaches a session-local tool registry.
|
||||
func (s *Server) SetToolRegistry(r *Registry, sessionID string) {
|
||||
s.toolRegistry = r
|
||||
s.sessionID = sessionID
|
||||
}
|
||||
|
||||
// callPromotedTool executes a tool promoted via the registry by running the
|
||||
// script file inside the sandbox, piping the JSON args to stdin.
|
||||
func (s *Server) callPromotedTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) {
|
||||
// Resolve script path.
|
||||
if strings.Contains(tool, "/") || strings.Contains(tool, "..") {
|
||||
return nil, fmt.Errorf("invalid tool name")
|
||||
}
|
||||
path, err := ResolveTool(tool)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Extract dispatch-level flags (not passed to tool script).
|
||||
bypassed := false
|
||||
timeout := 30
|
||||
sandboxName := "default"
|
||||
var argMap map[string]interface{}
|
||||
if err := json.Unmarshal(args, &argMap); err == nil {
|
||||
if e, ok := argMap["bypass"]; ok {
|
||||
switch v := e.(type) {
|
||||
case bool:
|
||||
bypassed = v
|
||||
case string:
|
||||
bypassed = v == "true" || v == "1"
|
||||
}
|
||||
delete(argMap, "bypass")
|
||||
}
|
||||
if t, ok := argMap["timeout"]; ok {
|
||||
switch v := t.(type) {
|
||||
case float64:
|
||||
timeout = int(v)
|
||||
case string:
|
||||
var n int
|
||||
if _, err := fmt.Sscanf(v, "%d", &n); err == nil {
|
||||
timeout = n
|
||||
}
|
||||
}
|
||||
delete(argMap, "timeout")
|
||||
}
|
||||
if s, ok := argMap["sandbox"]; ok {
|
||||
if v, ok := s.(string); ok && v != "" {
|
||||
sandboxName = v
|
||||
}
|
||||
delete(argMap, "sandbox")
|
||||
}
|
||||
args, _ = json.Marshal(argMap)
|
||||
}
|
||||
|
||||
// Check if tool requires sudo (from .meta).
|
||||
needsSudo := false
|
||||
if m, err := LoadMetaFile(tool); err == nil && m != nil {
|
||||
resolved := m.Resolve()
|
||||
if resolved != nil && resolved.Sudo {
|
||||
needsSudo = true
|
||||
bypassed = true // sudo implies bypass
|
||||
}
|
||||
}
|
||||
|
||||
// The tool script receives JSON args on stdin.
|
||||
code := path
|
||||
stdinData := string(args)
|
||||
|
||||
var result string
|
||||
if needsSudo {
|
||||
s.wdMu.RLock()
|
||||
workDir := s.cwd
|
||||
s.wdMu.RUnlock()
|
||||
sudoCode := fmt.Sprintf("cat <<'OLLIE_EOF' | %s\n%s\nOLLIE_EOF", code, stdinData)
|
||||
result, err = s.executeBypassSudo(ctx, sudoCode, workDir, timeout)
|
||||
} else if bypassed {
|
||||
s.wdMu.RLock()
|
||||
workDir := s.cwd
|
||||
s.wdMu.RUnlock()
|
||||
// Broker protocol has no stdin support; pipe JSON via heredoc.
|
||||
bypassCode := fmt.Sprintf("cat <<'OLLIE_EOF' | %s\n%s\nOLLIE_EOF", code, stdinData)
|
||||
result, err = s.executeBypass(ctx, bypassCode, workDir, timeout)
|
||||
} else {
|
||||
result, err = s.executeWithStdin(ctx, code, "bash", timeout, sandboxName, false, stdinData)
|
||||
}
|
||||
if err != nil {
|
||||
return json.Marshal(ToolResult{
|
||||
IsError: true,
|
||||
Content: []ToolResultContent{{Type: "text", Text: result + ": " + err.Error()}},
|
||||
})
|
||||
}
|
||||
return json.Marshal(ToolResult{
|
||||
Content: []ToolResultContent{{Type: "text", Text: result}},
|
||||
})
|
||||
}
|
||||
|
||||
// Close is called when the session ends. Calls OnClose hook if registered.
|
||||
func (s *Server) Close() {
|
||||
s.cleanupDetached()
|
||||
if s.OnClose != nil {
|
||||
s.OnClose()
|
||||
}
|
||||
}
|
||||
|
||||
// Detach signals the currently running process to be detached from the agent.
|
||||
// The process continues running; its output is captured in a ring buffer.
|
||||
// Returns false if no process is currently running.
|
||||
func (s *Server) Detach() bool {
|
||||
s.detachMu.Lock()
|
||||
ch := s.detachCh
|
||||
s.detachMu.Unlock()
|
||||
if ch == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case ch <- struct{}{}:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// ListDetachedRaw returns detached process info as []any (each element is map[string]any)
|
||||
// for consumption by packages that can't import this package directly.
|
||||
func (s *Server) ListDetachedRaw() []any {
|
||||
s.detachMu.Lock()
|
||||
defer s.detachMu.Unlock()
|
||||
out := make([]any, len(s.detached))
|
||||
for i, p := range s.detached {
|
||||
info := p.Info()
|
||||
out[i] = map[string]any{
|
||||
"pid": info.PID,
|
||||
"command": info.Command,
|
||||
"started": info.Started,
|
||||
"exited": info.Exited,
|
||||
"exit_code": info.ExitCode,
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// SignalDetached sends a signal to a detached process by PID.
|
||||
func (s *Server) SignalDetached(pid int, sig syscall.Signal) error {
|
||||
s.detachMu.Lock()
|
||||
defer s.detachMu.Unlock()
|
||||
for _, p := range s.detached {
|
||||
if p.PID == pid {
|
||||
p.Mu.Lock()
|
||||
defer p.Mu.Unlock()
|
||||
if p.Exited {
|
||||
return fmt.Errorf("process %d already exited", pid)
|
||||
}
|
||||
return syscall.Kill(-pid, sig)
|
||||
}
|
||||
}
|
||||
return fmt.Errorf("no detached process with pid %d", pid)
|
||||
}
|
||||
|
||||
// GetDetachedOutput returns the ring buffer contents for a detached process.
|
||||
func (s *Server) GetDetachedOutput(pid int) (string, error) {
|
||||
s.detachMu.Lock()
|
||||
defer s.detachMu.Unlock()
|
||||
for _, p := range s.detached {
|
||||
if p.PID == pid {
|
||||
return p.Output(), nil
|
||||
}
|
||||
}
|
||||
return "", fmt.Errorf("no detached process with pid %d", pid)
|
||||
}
|
||||
|
||||
// DismissDetached removes an exited process from the list.
|
||||
func (s *Server) DismissDetached(pid int) bool {
|
||||
s.detachMu.Lock()
|
||||
defer s.detachMu.Unlock()
|
||||
for i, p := range s.detached {
|
||||
if p.PID == pid && p.Exited {
|
||||
s.detached = append(s.detached[:i], s.detached[i+1:]...)
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// cleanupDetached sends SIGTERM to all running detached processes.
|
||||
func (s *Server) cleanupDetached() {
|
||||
for _, p := range s.detached {
|
||||
p.Mu.Lock()
|
||||
if !p.Exited {
|
||||
syscall.Kill(-p.PID, syscall.SIGTERM)
|
||||
}
|
||||
p.Mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
// SetOnToolsChanged sets the callback for tool directory changes.
|
||||
func (s *Server) SetOnToolsChanged(fn func()) {
|
||||
s.OnToolsChanged = fn
|
||||
}
|
||||
|
||||
// ToolRegistryRevision returns the current registry revision for this server's session.
|
||||
// Returns 0 if no registry is configured.
|
||||
func (s *Server) ToolRegistryRevision() uint64 {
|
||||
if s.toolRegistry == nil || s.sessionID == "" {
|
||||
return 0
|
||||
}
|
||||
return s.toolRegistry.Revision(s.sessionID)
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -1,86 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// universalPatterns apply to all code.
|
||||
var universalPatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`\bmkfs\b`),
|
||||
regexp.MustCompile(`\bdd\b.*\bif=/dev/`),
|
||||
regexp.MustCompile(`\b(sudo|su)\s`),
|
||||
regexp.MustCompile(`/etc/(shadow|sudoers)`),
|
||||
}
|
||||
|
||||
// bashPatterns apply to bash (flag syntax, redirects, shell-specific constructs).
|
||||
var bashPatterns = []*regexp.Regexp{
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*f[a-z]*\s+)*-[a-z]*r[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
|
||||
regexp.MustCompile(`rm\s+.*--recursive.*--force`),
|
||||
regexp.MustCompile(`rm\s+.*--force.*--recursive`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\.\.(/|\s|$)`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+~`),
|
||||
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\*`),
|
||||
regexp.MustCompile(`:\s*\(\s*\)\s*\{\s*:\s*\|\s*:\s*&`), // fork bomb
|
||||
regexp.MustCompile(`>\s*/dev/sd`),
|
||||
regexp.MustCompile(`\beval\s+".*\$`),
|
||||
}
|
||||
|
||||
var languagePatterns = map[string][]*regexp.Regexp{
|
||||
"bash": bashPatterns,
|
||||
"": bashPatterns,
|
||||
}
|
||||
|
||||
var whitespacePattern = regexp.MustCompile(`\s+`)
|
||||
|
||||
// ValidateCode checks code against dangerous patterns.
|
||||
func (s *Server) ValidateCode(code, language string) error {
|
||||
if err := s.checkRateLimit(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
normalized := strings.ToLower(code)
|
||||
normalized = whitespacePattern.ReplaceAllString(normalized, " ")
|
||||
|
||||
patterns := append(universalPatterns, languagePatterns[language]...)
|
||||
for _, pattern := range patterns {
|
||||
if pattern.MatchString(normalized) {
|
||||
s.recordValidationFailure()
|
||||
return fmt.Errorf("dangerous pattern detected")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) checkRateLimit() error {
|
||||
s.rateLimitMu.Lock()
|
||||
defer s.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if now.Before(s.blockedUntil) {
|
||||
remaining := s.blockedUntil.Sub(now).Round(time.Second)
|
||||
return &RateLimitedError{Remaining: remaining}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) recordValidationFailure() {
|
||||
s.rateLimitMu.Lock()
|
||||
defer s.rateLimitMu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if now.Sub(s.lastFailure) > failureWindow {
|
||||
s.validationFailures = 0
|
||||
}
|
||||
|
||||
s.validationFailures++
|
||||
s.lastFailure = now
|
||||
|
||||
if s.validationFailures >= maxFailures {
|
||||
s.blockedUntil = now.Add(blockDuration)
|
||||
s.validationFailures = 0
|
||||
}
|
||||
}
|
||||
|
|
@ -1,164 +0,0 @@
|
|||
package toolsrv_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"ollie/toolsrv"
|
||||
)
|
||||
|
||||
// TestSocketSpawnAndDial verifies the full lifecycle:
|
||||
// spawn ollie-remote → dial the socket → call RPCs → close.
|
||||
func TestSocketSpawnAndDial(t *testing.T) {
|
||||
// Skip if ollie-remote binary not available.
|
||||
if _, err := findBinary(); err != nil {
|
||||
t.Skipf("ollie-remote not found: %v", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
cwd, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Spawn the toolsrv subprocess.
|
||||
proc, err := toolsrv.Spawn(ctx, cwd, toolsrv.WithYolo())
|
||||
if err != nil {
|
||||
t.Fatalf("Spawn: %v", err)
|
||||
}
|
||||
defer proc.Close()
|
||||
|
||||
if proc.Socket == "" {
|
||||
t.Fatal("expected non-empty socket path")
|
||||
}
|
||||
|
||||
// Verify socket file exists.
|
||||
if _, err := os.Stat(proc.Socket); err != nil {
|
||||
t.Fatalf("socket file not found: %v", err)
|
||||
}
|
||||
|
||||
// Dial a connection.
|
||||
conn, err := toolsrv.Dial(ctx, proc.Socket)
|
||||
if err != nil {
|
||||
t.Fatalf("Dial: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Ping.
|
||||
if err := conn.Ping(); err != nil {
|
||||
t.Fatalf("Ping: %v", err)
|
||||
}
|
||||
|
||||
// FetchHostInfo.
|
||||
info, err := conn.FetchHostInfo()
|
||||
if err != nil {
|
||||
t.Fatalf("FetchHostInfo: %v", err)
|
||||
}
|
||||
if info.Platform == "" {
|
||||
t.Error("expected non-empty platform")
|
||||
}
|
||||
|
||||
// SetEnv (fire-and-forget, just verify no crash).
|
||||
conn.SetEnv("TEST_KEY", "test_value")
|
||||
|
||||
// SetCWD.
|
||||
conn.SetCWD("/tmp")
|
||||
|
||||
// ListTools. The remote server starts with no promoted tools; tools are
|
||||
// loaded dynamically through the session tool registry.
|
||||
if _, err := conn.ListTools(); err != nil {
|
||||
t.Fatalf("ListTools: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSocketMultipleClients verifies multiple clients can connect simultaneously.
|
||||
func TestSocketMultipleClients(t *testing.T) {
|
||||
if _, err := findBinary(); err != nil {
|
||||
t.Skipf("ollie-remote not found: %v", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
cwd, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
proc, err := toolsrv.Spawn(ctx, cwd, toolsrv.WithYolo())
|
||||
if err != nil {
|
||||
t.Fatalf("Spawn: %v", err)
|
||||
}
|
||||
defer proc.Close()
|
||||
|
||||
// Dial three clients concurrently.
|
||||
conns := make([]*toolsrv.Conn, 3)
|
||||
for i := range conns {
|
||||
c, err := toolsrv.Dial(ctx, proc.Socket)
|
||||
if err != nil {
|
||||
t.Fatalf("Dial[%d]: %v", i, err)
|
||||
}
|
||||
conns[i] = c
|
||||
}
|
||||
|
||||
// Each client can ping independently.
|
||||
for i, c := range conns {
|
||||
if err := c.Ping(); err != nil {
|
||||
t.Errorf("Ping[%d]: %v", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Close all.
|
||||
for _, c := range conns {
|
||||
c.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// TestSocketCleanup verifies the socket is removed after Close.
|
||||
func TestSocketCleanup(t *testing.T) {
|
||||
if _, err := findBinary(); err != nil {
|
||||
t.Skipf("ollie-remote not found: %v", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
cwd, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
proc, err := toolsrv.Spawn(ctx, cwd, toolsrv.WithYolo())
|
||||
if err != nil {
|
||||
t.Fatalf("Spawn: %v", err)
|
||||
}
|
||||
sockPath := proc.Socket
|
||||
proc.Close()
|
||||
|
||||
// Socket should be cleaned up.
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
if _, err := net.DialTimeout("unix", sockPath, time.Second); err == nil {
|
||||
t.Error("expected socket to be removed after Close")
|
||||
}
|
||||
}
|
||||
|
||||
// findBinary checks if ollie-remote is available.
|
||||
func findBinary() (string, error) {
|
||||
home, _ := os.UserHomeDir()
|
||||
candidates := []string{
|
||||
filepath.Join(home, "bin", "ollie-remote"),
|
||||
filepath.Join(home, ".local", "bin", "ollie-remote"),
|
||||
}
|
||||
for _, p := range candidates {
|
||||
if _, err := os.Stat(p); err == nil {
|
||||
return p, nil
|
||||
}
|
||||
}
|
||||
return "", os.ErrNotExist
|
||||
}
|
||||
|
|
@ -1,76 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// Process represents a running toolsrv subprocess.
|
||||
// The session owns this and calls Close() when the session ends.
|
||||
type Process struct {
|
||||
// Socket is the Unix socket path clients should Dial().
|
||||
Socket string
|
||||
// Info holds host environment details (platform, git repo status).
|
||||
Info HostInfo
|
||||
// cleanup shuts down the subprocess and removes the socket.
|
||||
cleanup func()
|
||||
}
|
||||
|
||||
// Close shuts down the toolsrv subprocess and removes the socket file.
|
||||
func (p *Process) Close() {
|
||||
if p.cleanup != nil {
|
||||
p.cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
// Spawn starts a local toolsrv subprocess listening on a Unix socket.
|
||||
// The returned Process owns the subprocess lifecycle. Callers connect
|
||||
// via Dial(LocalAddr(proc.Socket)).
|
||||
func Spawn(ctx context.Context, cwd string, opts ...Option) (*Process, error) {
|
||||
StreamOutput(ctx, "[toolsrv] spawning local toolsrv...\n")
|
||||
|
||||
t, err := localDial(ctx, cwd, opts...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Verify the server is responsive via the initial connection.
|
||||
conn := NewConnSplit(t.r, t.w, func() {})
|
||||
var info HostInfo
|
||||
if hi, err := conn.FetchHostInfo(); err == nil {
|
||||
info = hi
|
||||
}
|
||||
|
||||
StreamOutput(ctx, fmt.Sprintf("[toolsrv] local ready: %s (socket %s)\n", cwd, t.sockPath))
|
||||
|
||||
return &Process{
|
||||
Socket: t.sockPath,
|
||||
Info: info,
|
||||
cleanup: t.cleanup,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// SpawnRemote starts toolsrv on a remote host via SSH, sets up Unix
|
||||
// socket forwarding, and returns a Process with the local forwarded socket.
|
||||
func SpawnRemote(ctx context.Context, cfg RemoteConfig) (*Process, error) {
|
||||
StreamOutput(ctx, fmt.Sprintf("[toolsrv] dialing %s (cwd: %s)...\n", cfg.SSHTarget, cfg.CWD))
|
||||
|
||||
t, err := sshDial(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("spawn remote: %w", err)
|
||||
}
|
||||
|
||||
conn := NewConnSplit(t.r, t.w, func() {})
|
||||
var info HostInfo
|
||||
if hi, err := conn.FetchHostInfo(); err == nil {
|
||||
info = hi
|
||||
}
|
||||
|
||||
StreamOutput(ctx, fmt.Sprintf("[toolsrv] remote ready: %s → %s (%s/%s)\n", cfg.SSHTarget, t.sockPath, info.Platform, info.Arch))
|
||||
|
||||
return &Process{
|
||||
Socket: t.sockPath,
|
||||
Info: info,
|
||||
cleanup: t.cleanup,
|
||||
}, nil
|
||||
}
|
||||
|
|
@ -1,30 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import "context"
|
||||
|
||||
// OutputFunc is a callback for streaming partial tool output.
|
||||
type OutputFunc func(data string)
|
||||
|
||||
type streamKey struct{}
|
||||
|
||||
// WithOutputStream returns a context carrying an output streaming callback.
|
||||
// Tool servers can call StreamOutput(ctx, data) to emit partial results.
|
||||
func WithOutputStream(ctx context.Context, fn OutputFunc) context.Context {
|
||||
return context.WithValue(ctx, streamKey{}, fn)
|
||||
}
|
||||
|
||||
// StreamOutput sends partial output to the streaming callback in ctx, if any.
|
||||
func StreamOutput(ctx context.Context, data string) {
|
||||
if fn, ok := ctx.Value(streamKey{}).(OutputFunc); ok && fn != nil {
|
||||
fn(data)
|
||||
}
|
||||
}
|
||||
|
||||
// StreamFunc returns the streaming callback from ctx, or nil if none set.
|
||||
// Useful for passing to writers that call the function directly.
|
||||
func StreamFunc(ctx context.Context) func(string) {
|
||||
if fn, ok := ctx.Value(streamKey{}).(OutputFunc); ok {
|
||||
return func(data string) { fn(data) }
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
@ -1,287 +0,0 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"ollie/paths"
|
||||
)
|
||||
|
||||
// socketSeq is a monotonic counter to ensure unique socket paths.
|
||||
var socketSeq atomic.Int64
|
||||
|
||||
// transport encapsulates the plumbing for a JSON-RPC connection to a
|
||||
// toolsrv subprocess. Both local and SSH transports produce the same
|
||||
// result: a reader, writer, socket path, and cleanup function.
|
||||
type transport struct {
|
||||
r io.Reader
|
||||
w io.WriteCloser
|
||||
sockPath string
|
||||
cleanup func()
|
||||
}
|
||||
|
||||
// localDial spawns toolsrv as a local subprocess listening on a Unix
|
||||
// socket, waits for it to be ready, and returns the socket path.
|
||||
func localDial(ctx context.Context, cwd string, opts ...Option) (*transport, error) {
|
||||
// Bootstrap: verify toolsrv binary is present.
|
||||
binPath, err := findLocalBinary()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bootstrap: toolsrv not found")
|
||||
}
|
||||
|
||||
// Verify tools directory is accessible via ToolsPath().
|
||||
toolsDir := ToolsPath()
|
||||
if _, err := os.Stat(toolsDir); err != nil {
|
||||
return nil, fmt.Errorf("bootstrap: tools directory not found at %s", toolsDir)
|
||||
}
|
||||
|
||||
// Generate a unique socket path.
|
||||
sockDir := filepath.Join(paths.RuntimeDir(), "ollie")
|
||||
os.MkdirAll(sockDir, 0700)
|
||||
sockPath := filepath.Join(sockDir, fmt.Sprintf("toolsrv-%d-%d.sock", os.Getpid(), socketSeq.Add(1)))
|
||||
|
||||
args := []string{"serve", "--cwd", cwd, "--listen", sockPath}
|
||||
|
||||
probe := &Server{}
|
||||
for _, o := range opts {
|
||||
o(probe)
|
||||
}
|
||||
if probe.Yolo {
|
||||
args = append(args, "--yolo")
|
||||
}
|
||||
|
||||
cmd := exec.CommandContext(ctx, binPath, args...)
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{
|
||||
Setpgid: true,
|
||||
Pdeathsig: syscall.SIGTERM,
|
||||
}
|
||||
|
||||
// Inherit environment.
|
||||
cmd.Env = os.Environ()
|
||||
|
||||
// Watch stderr for readiness signal.
|
||||
stderrPipe, err := cmd.StderrPipe()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("local stderr pipe: %w", err)
|
||||
}
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
return nil, fmt.Errorf("local start: %w", err)
|
||||
}
|
||||
|
||||
// Wait for ListenReady on stderr.
|
||||
readyCh := make(chan struct{})
|
||||
go func() {
|
||||
scanner := bufio.NewScanner(stderrPipe)
|
||||
for scanner.Scan() {
|
||||
if strings.TrimSpace(scanner.Text()) == "ListenReady" {
|
||||
close(readyCh)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-readyCh:
|
||||
case <-time.After(10 * time.Second):
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGTERM)
|
||||
cmd.Wait()
|
||||
return nil, fmt.Errorf("local toolsrv startup timeout")
|
||||
}
|
||||
|
||||
// Connect to the socket.
|
||||
conn, err := net.Dial("unix", sockPath)
|
||||
if err != nil {
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGTERM)
|
||||
cmd.Wait()
|
||||
return nil, fmt.Errorf("local socket connect: %w", err)
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
conn.Close()
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGTERM)
|
||||
done := make(chan struct{})
|
||||
go func() { cmd.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
<-done
|
||||
}
|
||||
os.Remove(sockPath)
|
||||
}
|
||||
|
||||
return &transport{r: conn, w: conn, sockPath: sockPath, cleanup: cleanup}, nil
|
||||
}
|
||||
|
||||
// sshDial connects to a remote host via SSH, bootstraps toolsrv
|
||||
// with --listen mode, sets up SSH Unix socket forwarding, and returns
|
||||
// a transport connected to the forwarded local socket.
|
||||
//
|
||||
// Assumes pubkey auth (BatchMode=yes). No password handling.
|
||||
func sshDial(ctx context.Context, cfg RemoteConfig) (*transport, error) {
|
||||
// Bootstrap: verify local environment has everything needed.
|
||||
cacheDir := paths.CfgDir()
|
||||
if _, err := os.Stat(cacheDir); err != nil {
|
||||
return nil, fmt.Errorf("bootstrap: local environment not found at %s", cacheDir)
|
||||
}
|
||||
|
||||
// Generate socket paths.
|
||||
sockDir := filepath.Join(paths.RuntimeDir(), "ollie")
|
||||
os.MkdirAll(sockDir, 0700)
|
||||
localSock := filepath.Join(sockDir, fmt.Sprintf("remote-%s-%d-%d.sock", cfg.SSHTarget, os.Getpid(), socketSeq.Add(1)))
|
||||
remoteSock := fmt.Sprintf("/tmp/ollie-toolsrv-%d.sock", os.Getpid())
|
||||
|
||||
// Remote bootstrap: extract tarball, set env, exec serve.
|
||||
bootstrap := fmt.Sprintf(`#!/bin/sh
|
||||
set -e
|
||||
CACHE_DIR="${XDG_CACHE_HOME:-$HOME/.cache}/ollie"
|
||||
mkdir -p "$CACHE_DIR"
|
||||
tar xzf - -C "$CACHE_DIR"
|
||||
# CACHE_DIR is the extracted ollie config dir; set XDG_CONFIG_HOME so
|
||||
# $XDG_CONFIG_HOME/ollie resolves to it.
|
||||
export XDG_CONFIG_HOME="$(dirname "$CACHE_DIR")"
|
||||
export PATH="$CACHE_DIR/bin:$PATH"
|
||||
exec "$CACHE_DIR/bin/toolsrv" serve --cwd %s --listen %s
|
||||
`, shellEscape(cfg.CWD), shellEscape(remoteSock))
|
||||
|
||||
sshArgs := []string{
|
||||
"-o", "BatchMode=yes",
|
||||
"-o", "StrictHostKeyChecking=accept-new",
|
||||
"-o", "StreamLocalBindUnlink=yes",
|
||||
"-L", localSock + ":" + remoteSock,
|
||||
}
|
||||
sshTarget := cfg.SSHTarget
|
||||
if host, port, ok := strings.Cut(sshTarget, ":"); ok && port != "" {
|
||||
sshTarget = host
|
||||
sshArgs = append(sshArgs, "-p", port)
|
||||
}
|
||||
sshArgs = append(sshArgs, sshTarget, "bash -s")
|
||||
|
||||
cmd := exec.CommandContext(ctx, "ssh", sshArgs...)
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{
|
||||
Setpgid: true,
|
||||
Pdeathsig: syscall.SIGTERM,
|
||||
}
|
||||
|
||||
stdin, err := cmd.StdinPipe()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ssh stdin pipe: %w", err)
|
||||
}
|
||||
stderrPipe, err := cmd.StderrPipe()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ssh stderr pipe: %w", err)
|
||||
}
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
return nil, fmt.Errorf("ssh start: %w", err)
|
||||
}
|
||||
|
||||
// Send tarball over stdin, then the bootstrap script.
|
||||
if err := writeTarball(stdin, cacheDir); err != nil {
|
||||
cmd.Process.Kill()
|
||||
return nil, fmt.Errorf("send tarball: %w", err)
|
||||
}
|
||||
if _, err := stdin.Write([]byte(bootstrap + "\n")); err != nil {
|
||||
cmd.Process.Kill()
|
||||
return nil, fmt.Errorf("send bootstrap: %w", err)
|
||||
}
|
||||
|
||||
// Monitor stderr for ListenReady.
|
||||
readyCh := make(chan struct{}, 1)
|
||||
go func() {
|
||||
scanner := bufio.NewScanner(stderrPipe)
|
||||
for scanner.Scan() {
|
||||
if strings.TrimSpace(scanner.Text()) == "ListenReady" {
|
||||
close(readyCh)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-readyCh:
|
||||
case <-time.After(30 * time.Second):
|
||||
cmd.Process.Kill()
|
||||
return nil, fmt.Errorf("bootstrap timeout waiting for listen ready")
|
||||
}
|
||||
|
||||
// Poll for local socket forwarding.
|
||||
var conn net.Conn
|
||||
for i := 0; i < 20; i++ {
|
||||
conn, err = net.Dial("unix", localSock)
|
||||
if err == nil {
|
||||
break
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
if err != nil {
|
||||
cmd.Process.Kill()
|
||||
return nil, fmt.Errorf("ssh socket connect: %w", err)
|
||||
}
|
||||
|
||||
cleanup := func() {
|
||||
conn.Close()
|
||||
stdin.Close()
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGTERM)
|
||||
done := make(chan struct{})
|
||||
go func() { cmd.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(3 * time.Second):
|
||||
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
<-done
|
||||
}
|
||||
os.Remove(localSock)
|
||||
}
|
||||
|
||||
return &transport{r: conn, w: conn, sockPath: localSock, cleanup: cleanup}, nil
|
||||
}
|
||||
|
||||
// writeTarball creates a gzipped tarball of the directory at dir and writes
|
||||
// it to w. Used to transfer the ollie environment to remote hosts.
|
||||
func writeTarball(w io.Writer, dir string) error {
|
||||
cmd := exec.Command("tar", "czf", "-", "-C", dir, ".")
|
||||
cmd.Stdout = w
|
||||
cmd.Stderr = os.Stderr
|
||||
return cmd.Run()
|
||||
}
|
||||
|
||||
// findLocalBinary locates the toolsrv binary on the local system.
|
||||
func findLocalBinary() (string, error) {
|
||||
home, _ := os.UserHomeDir()
|
||||
candidates := []string{
|
||||
filepath.Join(home, "bin", "toolsrv"),
|
||||
filepath.Join(home, ".local", "bin", "toolsrv"),
|
||||
}
|
||||
for _, p := range candidates {
|
||||
if _, err := os.Stat(p); err == nil {
|
||||
return p, nil
|
||||
}
|
||||
}
|
||||
if p, err := exec.LookPath("toolsrv"); err == nil {
|
||||
return p, nil
|
||||
}
|
||||
return "", fmt.Errorf("toolsrv not found (checked %s)", strings.Join(candidates, ", "))
|
||||
}
|
||||
|
||||
// shellEscape escapes a string for safe use in a shell command.
|
||||
func shellEscape(s string) string {
|
||||
if s == "" {
|
||||
return "''"
|
||||
}
|
||||
if !strings.ContainsAny(s, " \t\n\r'\"\\$`!#&|;(){}") {
|
||||
return s
|
||||
}
|
||||
return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'"
|
||||
}
|
||||
Loading…
Reference in New Issue