ollie/cmd/toolsrv/internal/server/server.go

263 lines
7.4 KiB
Go

// spec.go - virtfs-based namespace specification for toolsrvclient.
package server
import (
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"fmt"
"runtime"
"strconv"
"strings"
"sync"
"ollie/cmd/toolsrv/internal/bypass"
"ollie/cmd/toolsrv/internal/registry"
"ollie/toolsrv/metadata"
"ollie/virtfs"
)
// Server holds the state for the toolsrv 9P server.
type Server struct {
mu sync.RWMutex
secret string // set on first auth, verified on subsequent
token string // session token returned after auth
registry *registry.Registry // tool registry
yolo bool // skip sandbox
// Process management
Fs *State
}
// NewServer creates a new toolsrv server.
func NewServer() *Server {
return &Server{
Fs: NewState(""), // cwd set per-agent
}
}
// SetRegistry configures the tool registry.
func (s *Server) SetRegistry(r *registry.Registry) {
s.mu.Lock()
s.registry = r
s.Fs.SetRegistry(r)
s.mu.Unlock()
}
// SetYolo enables/disables sandbox bypass.
func (s *Server) SetYolo(yolo bool) {
s.mu.Lock()
s.yolo = yolo
s.Fs.SetYolo(yolo)
s.mu.Unlock()
}
// Authenticate handles secret verification.
// First call sets the secret; subsequent calls must match it.
func (s *Server) Authenticate(clientSecret string) (string, error) {
s.mu.Lock()
defer s.mu.Unlock()
if s.secret == "" {
token, err := randomToken()
if err != nil {
return "", fmt.Errorf("generate authentication token: %w", err)
}
s.secret = clientSecret
s.token = token
return s.token, nil
}
if clientSecret != s.secret {
return "", fmt.Errorf("authentication failed")
}
return s.token, nil
}
// Token returns the current session token.
func (s *Server) Token() string {
s.mu.RLock()
defer s.mu.RUnlock()
return s.token
}
// HostInfo returns platform info as key=value lines.
func (s *Server) HostInfo() string {
return fmt.Sprintf("platform=%s\narch=%s\n", runtime.GOOS, runtime.GOARCH)
}
// BuildTree creates the virtfs tree for this server.
func (s *Server) BuildTree() *virtfs.Tree {
return virtfs.BuildTree(Spec(s))
}
// GenerateSecret generates a random shared secret for toolsrv auth.
func GenerateSecret() (string, error) {
secret := make([]byte, 32)
if _, err := rand.Read(secret); err != nil {
return "", err
}
return hex.EncodeToString(secret), nil
}
func randomToken() (string, error) {
b := make([]byte, 16)
if _, err := rand.Read(b); err != nil {
return "", err
}
return hex.EncodeToString(b), nil
}
// Spec returns the virtfs specification for the toolsrv namespace.
// The server is captured by closures — no context threading needed.
func Spec(srv *Server) virtfs.FsNodeDecl {
return virtfs.DirNode("/",
virtfs.FileNode("ctl", 0222,
virtfs.Doc("Control: load <agentID> <tool>, unload <agentID> <tool>, clear <agentID>"),
virtfs.Write(func(data []byte) error {
return srv.Fs.HandleCtl(strings.TrimSpace(string(data)))
}),
),
virtfs.FileNode("tools", 0666,
virtfs.Doc("Tools: write agent ID, read tool list (rdwr)"),
virtfs.Rdwr(func(_ context.Context, data []byte) ([]byte, error) {
agentID := strings.TrimSpace(string(data))
if agentID == "" {
return nil, fmt.Errorf("agent ID required")
}
return []byte(srv.Fs.HandleToolsRequest(agentID)), nil
}),
),
virtfs.FileNode("tools_rev", 0666,
virtfs.Doc("Tool registry revision: write agent ID, read revision"),
virtfs.Rdwr(func(_ context.Context, data []byte) ([]byte, error) {
return []byte(srv.Fs.HandleToolsRevision(string(data))), nil
}),
),
virtfs.FileNode("all", 0444,
virtfs.Doc("All available tools on disk: name<tab>description per line"),
virtfs.Read(func() ([]byte, error) {
tools := metadata.DiscoverTools()
var sb strings.Builder
for _, ti := range tools {
fmt.Fprintf(&sb, "%s\t%s\n", ti.Name, ti.Description)
}
return []byte(sb.String()), nil
}),
),
virtfs.FileNode("info", 0444,
virtfs.Doc("Host info: platform, arch"),
virtfs.Read(func() ([]byte, error) {
return []byte(srv.HostInfo()), nil
}),
),
virtfs.DirNode("bypass",
virtfs.FileNode("pending", 0444,
virtfs.Doc("Blocks until bypass request; returns JSON {id, cmd, cwd, env}"),
virtfs.BlockOnceRaw(func(ctx context.Context, _ string) ([]byte, string, error) {
req := bypass.NextPending(ctx)
if req == nil {
return nil, "", fmt.Errorf("context cancelled")
}
data, err := json.Marshal(req)
if err != nil {
return nil, "", err
}
return append(data, '\n'), "", nil
}),
),
virtfs.FileNode("resolve", 0222,
virtfs.Doc("Resolve bypass request: write JSON {id, approved} or {id, error}"),
virtfs.Write(func(data []byte) error {
var msg struct {
ID string `json:"id"`
Approved bool `json:"approved"`
Error string `json:"error"`
}
if err := json.Unmarshal(data, &msg); err != nil {
return fmt.Errorf("invalid JSON: %w", err)
}
if !bypass.Resolve(msg.ID, msg.Approved, msg.Error) {
return fmt.Errorf("unknown request ID: %s", msg.ID)
}
return nil
}),
),
),
virtfs.DirNode("proc",
virtfs.FileNode("list", 0666,
virtfs.Doc("List processes: write agent ID (or empty for all), read pid\\tstate\\ttool"),
virtfs.Rdwr(func(_ context.Context, data []byte) ([]byte, error) {
agentID := strings.TrimSpace(string(data))
return []byte(srv.Fs.ListProcsForAgent(agentID)), nil
}),
),
virtfs.FileNode("new", 0666,
virtfs.Doc("Execute tool: write token + tool + args, read result (blocking)"),
virtfs.Rdwr(func(ctx context.Context, data []byte) ([]byte, error) {
result, err := srv.Fs.HandleProcNew(ctx, string(data))
if err != nil {
return nil, err
}
return []byte(result), nil
}),
),
virtfs.FileNode("new.bg", 0666,
virtfs.Doc("Execute tool in background: write token + tool + args, read pid"),
virtfs.Rdwr(func(ctx context.Context, data []byte) ([]byte, error) {
result, err := srv.Fs.HandleProcNewBg(ctx, string(data))
if err != nil {
return nil, err
}
return []byte(result), nil
}),
),
virtfs.Each("{pid}", func() ([]virtfs.FsNodeDecl, error) {
pids := srv.Fs.ListProcs()
var out []virtfs.FsNodeDecl
for _, pid := range pids {
proc := srv.Fs.GetProc(pid)
if proc == nil {
continue
}
p := proc
out = append(out, virtfs.FsNodeDecl{
Name: strconv.Itoa(pid),
Children: []virtfs.FsNodeDecl{
virtfs.FileNode("out", 0444,
virtfs.Doc("Process output"),
virtfs.Read(func() ([]byte, error) {
return []byte(p.Output()), nil
}),
),
virtfs.FileNode("wait", 0444,
virtfs.Doc("Block until process exits, returns exit code"),
virtfs.BlockOnceRaw(func(_ context.Context, _ string) ([]byte, string, error) {
exitCode := p.Wait()
return []byte(fmt.Sprintf("%d\n", exitCode)), "", nil
}),
),
virtfs.FileNode("stat", 0444,
virtfs.Doc("Process status: running/exited, runtime, tool"),
virtfs.Read(func() ([]byte, error) {
return []byte(p.Stat()), nil
}),
),
virtfs.FileNode("ctl", 0222,
virtfs.Doc("Process control: write 'signal <N>' or 'dismiss'"),
virtfs.Write(func(data []byte) error {
return srv.Fs.HandleProcCtl(p.ID, strings.TrimSpace(string(data)))
}),
),
},
})
}
return out, nil
}),
),
)
}