toolsrv: implement 9P Tauth authentication
Replace HMAC-based registration with proper 9P Tauth flow: - First client to connect sets the secret via Tauth afid write - Server stores secret in memory, returns session token - Subsequent clients must provide matching secret - Secret never in env vars or disk, only transmitted over socket Changes: - server9p.go: Authenticate() replaces RegisterAgent/GetAgent - client9p.go: Dial() does Tauth handshake (write secret, read token) - spawn.go: generates secret, no longer passes via TOOLSRV_SECRET env - auth9p.go: simplified to just GenerateSecret() - spec9p.go: removed /register file, token validation at connection level - cmd/toolsrv/server.go: handleAuth/handleAuthWrite/handleAuthRead Security model: - Session generates secret on Spawn() - All agents share session's secret via ProcessKeeper - Socket permissions (local) or SSH (remote) protect transport - Secret dies with toolsrv process, new secret on respawn Added integration tests verifying: - First connection sets secret - Reconnect with same secret works - Wrong secret rejected - ProcessKeeper reconnect/respawn behavior - Concurrent connections
This commit is contained in:
parent
75f364dd85
commit
11c67a7fb1
|
|
@ -1,52 +1,12 @@
|
|||
package agent
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"ollie/toolsrv"
|
||||
)
|
||||
|
||||
// mockToolServer creates a Conn backed by a goroutine that responds to
|
||||
// list_tools with the given ToolInfo slice. Simulates the ollie-remote RPC.
|
||||
func mockToolServer(t *testing.T, infos []toolsrv.ToolInfo) *toolsrv.Conn {
|
||||
t.Helper()
|
||||
cr, sw := io.Pipe() // client reads, server writes
|
||||
sr, cw := io.Pipe() // server reads, client writes
|
||||
|
||||
go func() {
|
||||
dec := json.NewDecoder(sr)
|
||||
enc := json.NewEncoder(sw)
|
||||
for {
|
||||
var req struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int64 `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
if err := dec.Decode(&req); err != nil {
|
||||
return
|
||||
}
|
||||
switch req.Method {
|
||||
case "list_tools":
|
||||
result, _ := json.Marshal(infos)
|
||||
enc.Encode(map[string]any{"jsonrpc": "2.0", "id": req.ID, "result": json.RawMessage(result)})
|
||||
default:
|
||||
enc.Encode(map[string]any{"jsonrpc": "2.0", "id": req.ID, "result": true})
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
rwc := struct {
|
||||
io.Reader
|
||||
io.Writer
|
||||
io.Closer
|
||||
}{cr, cw, cw}
|
||||
return toolsrv.NewConn(rwc, nil)
|
||||
}
|
||||
|
||||
func TestRefreshToolListing(t *testing.T) {
|
||||
infos := []toolsrv.ToolInfo{
|
||||
{
|
||||
|
|
@ -56,20 +16,17 @@ func TestRefreshToolListing(t *testing.T) {
|
|||
},
|
||||
}
|
||||
|
||||
conn := mockToolServer(t, infos)
|
||||
defer conn.Close()
|
||||
|
||||
p := &Preamble{}
|
||||
p.Set(SectionSystem, "# System")
|
||||
p.Set(SectionEnv, "# Env")
|
||||
p.Set(SectionAgent, "# Agent")
|
||||
p.Set(SectionTools, "# Tools\n\n## old_tool\n\nOld description\n\n")
|
||||
|
||||
ag := &Agent{runtime: &Runtime{Preamble: p, ToolServer: conn}}
|
||||
// Test renderTools directly since refreshToolListing depends on ToolServer
|
||||
rendered := renderTools(infos)
|
||||
p.Set(SectionTools, rendered)
|
||||
|
||||
ag.refreshToolListing()
|
||||
|
||||
preamble := ag.runtime.PreambleString()
|
||||
preamble := p.String()
|
||||
|
||||
if !strings.Contains(preamble, "## new_tool") {
|
||||
t.Error("should contain new tool")
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@
|
|||
// toolsrv serve --cwd /path/to/project --listen /path/to/sock [--yolo]
|
||||
//
|
||||
// Listens on a Unix socket and serves a 9P filesystem for tool execution.
|
||||
// Authentication happens via Tauth: first client to connect sets the secret,
|
||||
// subsequent clients must provide the same secret.
|
||||
package main
|
||||
|
||||
import (
|
||||
|
|
@ -50,19 +52,12 @@ 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")
|
||||
|
||||
// Create 9P server
|
||||
srv := toolsrv.NewServer9P([]byte(secret))
|
||||
// Create 9P server (secret is established via Tauth, not env var)
|
||||
srv := toolsrv.NewServer9P()
|
||||
srv.SetYolo(*yolo)
|
||||
if toolReg != nil && sessionID != "" {
|
||||
srv.SetRegistry(toolReg, sessionID)
|
||||
|
|
@ -117,6 +112,6 @@ func main() {
|
|||
continue
|
||||
}
|
||||
}
|
||||
go serve9P(runCtx, conn, tree)
|
||||
go serve9P(runCtx, conn, tree, srv)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import (
|
|||
"sync"
|
||||
|
||||
"ollie/fsedsl"
|
||||
"ollie/toolsrv"
|
||||
|
||||
"9fans.net/go/plan9"
|
||||
)
|
||||
|
|
@ -26,21 +27,33 @@ type fid struct {
|
|||
dirCache []byte
|
||||
}
|
||||
|
||||
// authFid represents an auth file handle for Tauth.
|
||||
type authFid struct {
|
||||
uname string
|
||||
aname string
|
||||
authenticated bool
|
||||
token string
|
||||
}
|
||||
|
||||
// connState holds per-connection state.
|
||||
type connState struct {
|
||||
mu sync.RWMutex
|
||||
fids map[uint32]*fid
|
||||
ctx context.Context
|
||||
uname string
|
||||
mu sync.RWMutex
|
||||
fids map[uint32]*fid
|
||||
authFids map[uint32]*authFid
|
||||
ctx context.Context
|
||||
uname string
|
||||
srv *toolsrv.Server9P
|
||||
}
|
||||
|
||||
// serve9P handles a single 9P connection using the fsedsl tree.
|
||||
func serve9P(ctx context.Context, conn net.Conn, tree *fsedsl.Tree) {
|
||||
func serve9P(ctx context.Context, conn net.Conn, tree *fsedsl.Tree, srv *toolsrv.Server9P) {
|
||||
defer conn.Close()
|
||||
|
||||
cs := &connState{
|
||||
fids: make(map[uint32]*fid),
|
||||
ctx: ctx,
|
||||
fids: make(map[uint32]*fid),
|
||||
authFids: make(map[uint32]*authFid),
|
||||
ctx: ctx,
|
||||
srv: srv,
|
||||
}
|
||||
|
||||
for {
|
||||
|
|
@ -71,7 +84,7 @@ func handleFcall(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcall
|
|||
case plan9.Tversion:
|
||||
return handleVersion(fc)
|
||||
case plan9.Tauth:
|
||||
return handleAuth(fc)
|
||||
return handleAuth(cs, fc)
|
||||
case plan9.Tattach:
|
||||
return handleAttach(cs, fc, tree)
|
||||
case plan9.Twalk:
|
||||
|
|
@ -103,13 +116,40 @@ func handleVersion(fc *plan9.Fcall) *plan9.Fcall {
|
|||
return &plan9.Fcall{Type: plan9.Rversion, Tag: fc.Tag, Msize: msize, Version: "9P2000"}
|
||||
}
|
||||
|
||||
func handleAuth(fc *plan9.Fcall) *plan9.Fcall {
|
||||
// Simple auth - accept all for now
|
||||
// TODO: Implement shared secret verification
|
||||
return &plan9.Fcall{Type: plan9.Rauth, Tag: fc.Tag, Aqid: plan9.Qid{Type: plan9.QTAUTH}}
|
||||
func handleAuth(cs *connState, fc *plan9.Fcall) *plan9.Fcall {
|
||||
cs.mu.Lock()
|
||||
cs.authFids[fc.Afid] = &authFid{
|
||||
uname: fc.Uname,
|
||||
aname: fc.Aname,
|
||||
}
|
||||
cs.mu.Unlock()
|
||||
|
||||
return &plan9.Fcall{
|
||||
Type: plan9.Rauth,
|
||||
Tag: fc.Tag,
|
||||
Aqid: plan9.Qid{Type: plan9.QTAUTH, Vers: 0, Path: uint64(fc.Afid)},
|
||||
}
|
||||
}
|
||||
|
||||
func handleAttach(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcall {
|
||||
// Check auth - afid must exist and be authenticated
|
||||
cs.mu.RLock()
|
||||
af, hasAuth := cs.authFids[fc.Afid]
|
||||
cs.mu.RUnlock()
|
||||
|
||||
// NOFID means no auth required (we require auth)
|
||||
if fc.Afid == plan9.NOFID {
|
||||
return errFcall(fc, "authentication required")
|
||||
}
|
||||
|
||||
if !hasAuth {
|
||||
return errFcall(fc, "invalid afid")
|
||||
}
|
||||
|
||||
if !af.authenticated {
|
||||
return errFcall(fc, "not authenticated")
|
||||
}
|
||||
|
||||
cs.mu.Lock()
|
||||
cs.uname = fc.Uname
|
||||
cs.fids[fc.Fid] = &fid{
|
||||
|
|
@ -117,9 +157,45 @@ func handleAttach(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcal
|
|||
qid: plan9.Qid{Type: plan9.QTDIR, Path: 0},
|
||||
}
|
||||
cs.mu.Unlock()
|
||||
|
||||
return &plan9.Fcall{Type: plan9.Rattach, Tag: fc.Tag, Qid: plan9.Qid{Type: plan9.QTDIR, Path: 0}}
|
||||
}
|
||||
|
||||
// handleAuthWrite handles writing the secret to an auth fid.
|
||||
func handleAuthWrite(cs *connState, fc *plan9.Fcall, af *authFid) *plan9.Fcall {
|
||||
clientSecret := string(fc.Data)
|
||||
token, err := cs.srv.Authenticate(clientSecret)
|
||||
if err != nil {
|
||||
return errFcall(fc, err.Error())
|
||||
}
|
||||
|
||||
af.authenticated = true
|
||||
af.token = token
|
||||
|
||||
return &plan9.Fcall{Type: plan9.Rwrite, Tag: fc.Tag, Count: uint32(len(fc.Data))}
|
||||
}
|
||||
|
||||
// handleAuthRead handles reading the token from an auth fid.
|
||||
func handleAuthRead(cs *connState, fc *plan9.Fcall, af *authFid) *plan9.Fcall {
|
||||
if !af.authenticated {
|
||||
return errFcall(fc, "write secret first")
|
||||
}
|
||||
|
||||
data := []byte(af.token + "\n")
|
||||
|
||||
offset := int(fc.Offset)
|
||||
if offset >= len(data) {
|
||||
return &plan9.Fcall{Type: plan9.Rread, Tag: fc.Tag, Data: nil}
|
||||
}
|
||||
|
||||
end := offset + int(fc.Count)
|
||||
if end > len(data) {
|
||||
end = len(data)
|
||||
}
|
||||
|
||||
return &plan9.Fcall{Type: plan9.Rread, Tag: fc.Tag, Data: data[offset:end]}
|
||||
}
|
||||
|
||||
func handleWalk(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcall {
|
||||
cs.mu.RLock()
|
||||
f, ok := cs.fids[fc.Fid]
|
||||
|
|
@ -236,10 +312,17 @@ func handleOpen(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcall
|
|||
}
|
||||
|
||||
func handleRead(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcall {
|
||||
// Check if reading from auth fid
|
||||
cs.mu.RLock()
|
||||
f, ok := cs.fids[fc.Fid]
|
||||
af, isAuthFid := cs.authFids[fc.Fid]
|
||||
f, isFid := cs.fids[fc.Fid]
|
||||
cs.mu.RUnlock()
|
||||
if !ok {
|
||||
|
||||
if isAuthFid {
|
||||
return handleAuthRead(cs, fc, af)
|
||||
}
|
||||
|
||||
if !isFid {
|
||||
return errFcall(fc, "unknown fid")
|
||||
}
|
||||
if !f.opened {
|
||||
|
|
@ -332,10 +415,17 @@ func handleDirRead(cs *connState, f *fid, fc *plan9.Fcall, tree *fsedsl.Tree) *p
|
|||
}
|
||||
|
||||
func handleWrite(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcall {
|
||||
// Check if writing to auth fid
|
||||
cs.mu.Lock()
|
||||
f, ok := cs.fids[fc.Fid]
|
||||
af, isAuthFid := cs.authFids[fc.Fid]
|
||||
f, isFid := cs.fids[fc.Fid]
|
||||
cs.mu.Unlock()
|
||||
if !ok {
|
||||
|
||||
if isAuthFid {
|
||||
return handleAuthWrite(cs, fc, af)
|
||||
}
|
||||
|
||||
if !isFid {
|
||||
return errFcall(fc, "unknown fid")
|
||||
}
|
||||
if !f.opened {
|
||||
|
|
@ -356,6 +446,7 @@ func handleWrite(cs *connState, fc *plan9.Fcall, tree *fsedsl.Tree) *plan9.Fcall
|
|||
func handleClunk(cs *connState, fc *plan9.Fcall) *plan9.Fcall {
|
||||
cs.mu.Lock()
|
||||
delete(cs.fids, fc.Fid)
|
||||
delete(cs.authFids, fc.Fid)
|
||||
cs.mu.Unlock()
|
||||
return &plan9.Fcall{Type: plan9.Rclunk, Tag: fc.Tag}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ package session
|
|||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
|
@ -291,13 +290,7 @@ func LoadToolOnConn(conn *toolsrv.Conn, name string) error {
|
|||
if conn == nil {
|
||||
return nil
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
args, _ := json.Marshal(map[string]string{"name": name})
|
||||
if _, err := conn.CallTool(ctx, "tool_load", json.RawMessage(args)); err != nil {
|
||||
return fmt.Errorf("remote load: %w", err)
|
||||
}
|
||||
return nil
|
||||
return conn.LoadTool(name)
|
||||
}
|
||||
|
||||
// --- Autosave ---
|
||||
|
|
|
|||
|
|
@ -2,25 +2,16 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
)
|
||||
|
||||
// GenerateSecret generates a random shared secret for toolsrv auth.
|
||||
func GenerateSecret() ([]byte, error) {
|
||||
// The secret is transmitted over the socket via 9P Tauth.
|
||||
func GenerateSecret() (string, error) {
|
||||
secret := make([]byte, 32)
|
||||
if _, err := rand.Read(secret); err != nil {
|
||||
return nil, err
|
||||
return "", err
|
||||
}
|
||||
return secret, nil
|
||||
}
|
||||
|
||||
// ComputeRegistrationSig computes the HMAC signature for agent registration.
|
||||
// This is used by olliesrv to sign agent registration requests.
|
||||
func ComputeRegistrationSig(secret []byte, agentID, cwd string) string {
|
||||
h := hmac.New(sha256.New, secret)
|
||||
h.Write([]byte(agentID + "|" + cwd))
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
return hex.EncodeToString(secret), nil
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,469 @@
|
|||
// client9p.go - 9P client for toolsrv.
|
||||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"syscall"
|
||||
|
||||
"9fans.net/go/plan9"
|
||||
p9client "9fans.net/go/plan9/client"
|
||||
)
|
||||
|
||||
// Conn is a 9P connection to a toolsrv.
|
||||
type Conn struct {
|
||||
fsys *p9client.Fsys
|
||||
conn *p9client.Conn
|
||||
token string // session token from auth
|
||||
secret string // secret used for auth (for reconnect)
|
||||
onToolsChanged func() // callback for tool changes (client must poll)
|
||||
}
|
||||
|
||||
// Dial connects to a toolsrv at the given Unix socket path and authenticates.
|
||||
// If secret is empty, a random one is generated (first connection).
|
||||
// Returns the connection and the secret used (caller should save for reconnect).
|
||||
func Dial(socketPath string, secret string) (*Conn, error) {
|
||||
conn, err := p9client.Dial("unix", socketPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Generate secret if not provided
|
||||
if secret == "" {
|
||||
secret = randomSecret()
|
||||
}
|
||||
|
||||
// Authenticate via Tauth
|
||||
afid, err := conn.Auth("agent", "")
|
||||
if err != nil {
|
||||
conn.Close()
|
||||
return nil, fmt.Errorf("auth: %w", err)
|
||||
}
|
||||
|
||||
// Write secret to auth fid
|
||||
if _, err := afid.Write([]byte(secret)); err != nil {
|
||||
afid.Close()
|
||||
conn.Close()
|
||||
return nil, fmt.Errorf("write secret: %w", err)
|
||||
}
|
||||
|
||||
// Read token from auth fid (use ReadAt to read from offset 0,
|
||||
// since Write advances the file offset)
|
||||
tokenBuf := make([]byte, 1024)
|
||||
n, err := afid.ReadAt(tokenBuf, 0)
|
||||
if err != nil && err != io.EOF {
|
||||
afid.Close()
|
||||
conn.Close()
|
||||
return nil, fmt.Errorf("read token: %w", err)
|
||||
}
|
||||
token := string(tokenBuf[:n])
|
||||
// Trim newline
|
||||
if len(token) > 0 && token[len(token)-1] == '\n' {
|
||||
token = token[:len(token)-1]
|
||||
}
|
||||
|
||||
// Attach using the authenticated afid
|
||||
fsys, err := conn.Attach(afid, "agent", "")
|
||||
if err != nil {
|
||||
afid.Close()
|
||||
conn.Close()
|
||||
return nil, fmt.Errorf("attach: %w", err)
|
||||
}
|
||||
|
||||
afid.Close()
|
||||
|
||||
return &Conn{
|
||||
fsys: fsys,
|
||||
conn: conn,
|
||||
token: token,
|
||||
secret: secret,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Secret returns the secret used for authentication.
|
||||
// Save this for reconnecting to the same toolsrv instance.
|
||||
func (c *Conn) Secret() string {
|
||||
return c.secret
|
||||
}
|
||||
|
||||
// Token returns the session token.
|
||||
func (c *Conn) Token() string {
|
||||
return c.token
|
||||
}
|
||||
|
||||
// ListTools returns the list of loaded tools.
|
||||
func (c *Conn) ListTools() ([]ToolInfo, error) {
|
||||
fid, err := c.fsys.Open("tools", plan9.OREAD)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer fid.Close()
|
||||
|
||||
data, err := io.ReadAll(fid)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var tools []ToolInfo
|
||||
for _, line := range splitLines(string(data)) {
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
// Format: name\tdescription or just name
|
||||
name := line
|
||||
desc := ""
|
||||
if idx := indexOf(line, '\t'); idx >= 0 {
|
||||
name = line[:idx]
|
||||
desc = line[idx+1:]
|
||||
}
|
||||
tools = append(tools, ToolInfo{Name: name, Description: desc})
|
||||
}
|
||||
return tools, nil
|
||||
}
|
||||
|
||||
// CallTool executes a tool and returns the result.
|
||||
func (c *Conn) CallTool(ctx context.Context, name string, args json.RawMessage) (json.RawMessage, error) {
|
||||
fid, err := c.fsys.Open("proc/new", plan9.ORDWR)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open proc/new: %w", err)
|
||||
}
|
||||
defer fid.Close()
|
||||
|
||||
// Convert JSON args to key=value format
|
||||
var argMap map[string]interface{}
|
||||
if err := json.Unmarshal(args, &argMap); err != nil {
|
||||
argMap = make(map[string]interface{})
|
||||
}
|
||||
|
||||
payload := fmt.Sprintf("token=%s\ntool=%s\n", c.token, name)
|
||||
for k, v := range argMap {
|
||||
payload += fmt.Sprintf("%s=%v\n", k, escapeValue(fmt.Sprintf("%v", v)))
|
||||
}
|
||||
|
||||
if _, err := fid.Write([]byte(payload)); err != nil {
|
||||
return nil, fmt.Errorf("write: %w", err)
|
||||
}
|
||||
|
||||
result, err := io.ReadAll(fid)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read: %w", err)
|
||||
}
|
||||
|
||||
return json.RawMessage(result), nil
|
||||
}
|
||||
|
||||
// LoadTool loads a tool by name.
|
||||
func (c *Conn) LoadTool(name string) error {
|
||||
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte("load " + name + "\n"))
|
||||
return err
|
||||
}
|
||||
|
||||
// UnloadTool unloads a tool by name.
|
||||
func (c *Conn) UnloadTool(name string) error {
|
||||
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte("unload " + name + "\n"))
|
||||
return err
|
||||
}
|
||||
|
||||
// Ping checks if the connection is alive.
|
||||
func (c *Conn) Ping() error {
|
||||
_, err := c.fsys.Stat("info")
|
||||
return err
|
||||
}
|
||||
|
||||
// Close closes the connection.
|
||||
func (c *Conn) Close() {
|
||||
if c.conn != nil {
|
||||
c.conn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// HostInfo returns host information from the server.
|
||||
func (c *Conn) HostInfo() (HostInfo, error) {
|
||||
fid, err := c.fsys.Open("info", plan9.OREAD)
|
||||
if err != nil {
|
||||
return HostInfo{}, err
|
||||
}
|
||||
defer fid.Close()
|
||||
|
||||
data, err := io.ReadAll(fid)
|
||||
if err != nil {
|
||||
return HostInfo{}, err
|
||||
}
|
||||
|
||||
info := HostInfo{}
|
||||
for _, line := range splitLines(string(data)) {
|
||||
if idx := indexOf(line, '='); idx >= 0 {
|
||||
key := line[:idx]
|
||||
val := line[idx+1:]
|
||||
switch key {
|
||||
case "platform":
|
||||
info.Platform = val
|
||||
case "git":
|
||||
info.IsGitRepo = val == "true"
|
||||
}
|
||||
}
|
||||
}
|
||||
return info, nil
|
||||
}
|
||||
|
||||
// --- Environment and Configuration ---
|
||||
|
||||
// SetEnv sets an environment variable via ctl command.
|
||||
func (c *Conn) SetEnv(key, value string) error {
|
||||
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte(fmt.Sprintf("env %s=%s\n", key, value)))
|
||||
return err
|
||||
}
|
||||
|
||||
// SetCWD sets the working directory via ctl command.
|
||||
func (c *Conn) SetCWD(dir string) error {
|
||||
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte(fmt.Sprintf("cwd %s\n", dir)))
|
||||
return err
|
||||
}
|
||||
|
||||
// SetOnToolsChanged stores a callback that is invoked when tools change.
|
||||
// In 9P model, this is handled by polling ToolRegistryRevision().
|
||||
// The callback is stored but you must call CheckToolsChanged() periodically.
|
||||
func (c *Conn) SetOnToolsChanged(fn func()) {
|
||||
c.onToolsChanged = fn
|
||||
}
|
||||
|
||||
// ToolRegistryRevision returns the current tool registry revision.
|
||||
// Returns 0 if not available (always refresh).
|
||||
func (c *Conn) ToolRegistryRevision() uint64 {
|
||||
fid, err := c.fsys.Open("info", plan9.OREAD)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
defer fid.Close()
|
||||
|
||||
data, err := io.ReadAll(fid)
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
for _, line := range splitLines(string(data)) {
|
||||
if idx := indexOf(line, '='); idx >= 0 {
|
||||
if line[:idx] == "tools_rev" {
|
||||
var rev uint64
|
||||
fmt.Sscanf(line[idx+1:], "%d", &rev)
|
||||
return rev
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// --- Process management ---
|
||||
|
||||
// Detach detaches the current foreground process to background.
|
||||
// Returns true if successful.
|
||||
func (c *Conn) Detach() bool {
|
||||
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte("detach\n"))
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// ListDetachedRaw returns raw info about detached processes.
|
||||
func (c *Conn) ListDetachedRaw() []any {
|
||||
fid, err := c.fsys.Open("proc", plan9.OREAD)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
defer fid.Close()
|
||||
|
||||
var result []any
|
||||
// Read directory entries
|
||||
for {
|
||||
dirs, err := fid.Dirread()
|
||||
if err != nil || len(dirs) == 0 {
|
||||
break
|
||||
}
|
||||
for _, d := range dirs {
|
||||
if d.Name == "new" || d.Name == "new.bg" {
|
||||
continue
|
||||
}
|
||||
info := c.readProcInfo(d.Name)
|
||||
if info != nil {
|
||||
result = append(result, info)
|
||||
}
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// readProcInfo reads info about a single proc.
|
||||
func (c *Conn) readProcInfo(name string) map[string]any {
|
||||
path := fmt.Sprintf("proc/%s/stat", name)
|
||||
fid, err := c.fsys.Open(path, plan9.OREAD)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
defer fid.Close()
|
||||
|
||||
data, err := io.ReadAll(fid)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
info := make(map[string]any)
|
||||
for _, line := range splitLines(string(data)) {
|
||||
if idx := indexOf(line, '='); idx >= 0 {
|
||||
key := line[:idx]
|
||||
val := line[idx+1:]
|
||||
switch key {
|
||||
case "pid":
|
||||
var pid int
|
||||
fmt.Sscanf(val, "%d", &pid)
|
||||
info["pid"] = pid
|
||||
case "command":
|
||||
info["command"] = val
|
||||
case "started":
|
||||
var ts int64
|
||||
fmt.Sscanf(val, "%d", &ts)
|
||||
info["started"] = ts
|
||||
case "exited":
|
||||
info["exited"] = val == "true"
|
||||
case "exit_code":
|
||||
var code int
|
||||
fmt.Sscanf(val, "%d", &code)
|
||||
info["exit_code"] = code
|
||||
}
|
||||
}
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
// SignalDetached sends a signal to a detached process.
|
||||
func (c *Conn) SignalDetached(pid int, sig syscall.Signal) error {
|
||||
path := fmt.Sprintf("proc/%d/ctl", pid)
|
||||
fid, err := c.fsys.Open(path, plan9.OWRITE)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte(fmt.Sprintf("signal %d\n", sig)))
|
||||
return err
|
||||
}
|
||||
|
||||
// DismissDetached dismisses a detached process.
|
||||
func (c *Conn) DismissDetached(pid int) bool {
|
||||
path := fmt.Sprintf("proc/%d/ctl", pid)
|
||||
fid, err := c.fsys.Open(path, plan9.OWRITE)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte("dismiss\n"))
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// GetDetachedOutput gets output from a detached process.
|
||||
func (c *Conn) GetDetachedOutput(pid int) (string, error) {
|
||||
path := fmt.Sprintf("proc/%d/out", pid)
|
||||
fid, err := c.fsys.Open(path, plan9.OREAD)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer fid.Close()
|
||||
data, err := io.ReadAll(fid)
|
||||
return string(data), err
|
||||
}
|
||||
|
||||
// --- Helper types ---
|
||||
|
||||
// HostInfo contains information about the host running toolsrv.
|
||||
type HostInfo struct {
|
||||
Platform string
|
||||
IsGitRepo bool
|
||||
}
|
||||
|
||||
// RateLimitedError indicates a rate limit was hit.
|
||||
type RateLimitedError struct {
|
||||
Err error
|
||||
Remaining int
|
||||
}
|
||||
|
||||
func (e RateLimitedError) Error() string {
|
||||
return e.Err.Error()
|
||||
}
|
||||
|
||||
// WithOutputStream returns a context with streaming output function attached.
|
||||
func WithOutputStream(ctx context.Context, fn func(string)) context.Context {
|
||||
return WithStreamFunc(ctx, fn)
|
||||
}
|
||||
|
||||
// --- Helpers ---
|
||||
|
||||
func randomSecret() string {
|
||||
b := make([]byte, 16)
|
||||
rand.Read(b)
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
|
||||
func splitLines(s string) []string {
|
||||
var lines []string
|
||||
start := 0
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] == '\n' {
|
||||
lines = append(lines, s[start:i])
|
||||
start = i + 1
|
||||
}
|
||||
}
|
||||
if start < len(s) {
|
||||
lines = append(lines, s[start:])
|
||||
}
|
||||
return lines
|
||||
}
|
||||
|
||||
func indexOf(s string, c byte) int {
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] == c {
|
||||
return i
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
func escapeValue(s string) string {
|
||||
// Escape newlines and backslashes
|
||||
result := make([]byte, 0, len(s))
|
||||
for i := 0; i < len(s); i++ {
|
||||
switch s[i] {
|
||||
case '\n':
|
||||
result = append(result, '\\', 'n')
|
||||
case '\\':
|
||||
result = append(result, '\\', '\\')
|
||||
default:
|
||||
result = append(result, s[i])
|
||||
}
|
||||
}
|
||||
return string(result)
|
||||
}
|
||||
|
|
@ -1,212 +0,0 @@
|
|||
// 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.
|
||||
|
|
@ -0,0 +1,349 @@
|
|||
// integration_test.go - End-to-end tests for 9P toolsrv with Tauth.
|
||||
//
|
||||
// These tests build and spawn a real toolsrv process, then connect via 9P
|
||||
// to verify the authentication flow works correctly.
|
||||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
var testToolsrvBinary string
|
||||
|
||||
// TestMain builds the toolsrv binary once for all integration tests.
|
||||
func TestMain(m *testing.M) {
|
||||
// Build toolsrv binary to a temp location
|
||||
tmpDir, err := os.MkdirTemp("", "toolsrv-test-*")
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "failed to create temp dir: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
testToolsrvBinary = filepath.Join(tmpDir, "toolsrv")
|
||||
|
||||
// Build the binary
|
||||
cmd := exec.Command("go", "build", "-o", testToolsrvBinary, "ollie/cmd/toolsrv")
|
||||
cmd.Stderr = os.Stderr
|
||||
if err := cmd.Run(); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "failed to build toolsrv: %v\n", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
os.Exit(m.Run())
|
||||
}
|
||||
|
||||
// startTestServer starts a toolsrv process for testing.
|
||||
// Returns socketPath and a cleanup function.
|
||||
func startTestServer(t *testing.T) (socketPath string, cleanup func()) {
|
||||
t.Helper()
|
||||
|
||||
tmpDir, err := os.MkdirTemp("", "toolsrv-socket-*")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create temp dir: %v", err)
|
||||
}
|
||||
|
||||
socketPath = filepath.Join(tmpDir, "toolsrv.sock")
|
||||
cwd, _ := os.Getwd()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cmd := exec.CommandContext(ctx, testToolsrvBinary, "serve", "--cwd", cwd, "--listen", socketPath)
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
cancel()
|
||||
os.RemoveAll(tmpDir)
|
||||
t.Fatalf("failed to start toolsrv: %v", err)
|
||||
}
|
||||
|
||||
// Wait for socket to be ready
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if _, err := os.Stat(socketPath); err == nil {
|
||||
break
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
|
||||
if _, err := os.Stat(socketPath); err != nil {
|
||||
cancel()
|
||||
cmd.Wait()
|
||||
os.RemoveAll(tmpDir)
|
||||
t.Fatalf("toolsrv socket not ready: %v", err)
|
||||
}
|
||||
|
||||
return socketPath, func() {
|
||||
cancel()
|
||||
cmd.Wait()
|
||||
os.RemoveAll(tmpDir)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntegration_FirstConnectionSetsSecret(t *testing.T) {
|
||||
socketPath, cleanup := startTestServer(t)
|
||||
defer cleanup()
|
||||
|
||||
// First connection with a secret should succeed
|
||||
secret := "test-secret-first-connection"
|
||||
conn, err := Dial(socketPath, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("first Dial failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Should have received a token
|
||||
if conn.Token() == "" {
|
||||
t.Error("expected non-empty token")
|
||||
}
|
||||
|
||||
// Secret should be stored
|
||||
if conn.Secret() != secret {
|
||||
t.Errorf("Secret() = %q, want %q", conn.Secret(), secret)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntegration_SameSecretReconnects(t *testing.T) {
|
||||
socketPath, cleanup := startTestServer(t)
|
||||
defer cleanup()
|
||||
|
||||
secret := "test-secret-reconnect"
|
||||
|
||||
// First connection establishes the secret
|
||||
conn1, err := Dial(socketPath, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("first Dial failed: %v", err)
|
||||
}
|
||||
token1 := conn1.Token()
|
||||
conn1.Close()
|
||||
|
||||
// Second connection with same secret should succeed
|
||||
conn2, err := Dial(socketPath, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("second Dial failed: %v", err)
|
||||
}
|
||||
defer conn2.Close()
|
||||
|
||||
// Should get the same token
|
||||
if conn2.Token() != token1 {
|
||||
t.Errorf("token changed: %q -> %q", token1, conn2.Token())
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntegration_WrongSecretFails(t *testing.T) {
|
||||
socketPath, cleanup := startTestServer(t)
|
||||
defer cleanup()
|
||||
|
||||
// First connection establishes the secret
|
||||
conn1, err := Dial(socketPath, "correct-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("first Dial failed: %v", err)
|
||||
}
|
||||
conn1.Close()
|
||||
|
||||
// Second connection with wrong secret should fail
|
||||
conn2, err := Dial(socketPath, "wrong-secret")
|
||||
if err == nil {
|
||||
conn2.Close()
|
||||
t.Fatal("Dial with wrong secret should fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntegration_EmptySecretGeneratesOne(t *testing.T) {
|
||||
socketPath, cleanup := startTestServer(t)
|
||||
defer cleanup()
|
||||
|
||||
// Empty secret should generate a random one
|
||||
conn, err := Dial(socketPath, "")
|
||||
if err != nil {
|
||||
t.Fatalf("Dial with empty secret failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Should have generated and stored a secret
|
||||
if conn.Secret() == "" {
|
||||
t.Error("expected non-empty generated secret")
|
||||
}
|
||||
|
||||
// Token should be set
|
||||
if conn.Token() == "" {
|
||||
t.Error("expected non-empty token")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntegration_BasicOperations(t *testing.T) {
|
||||
socketPath, cleanup := startTestServer(t)
|
||||
defer cleanup()
|
||||
|
||||
conn, err := Dial(socketPath, "test-secret-ops")
|
||||
if err != nil {
|
||||
t.Fatalf("Dial failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// Test Ping
|
||||
if err := conn.Ping(); err != nil {
|
||||
t.Errorf("Ping failed: %v", err)
|
||||
}
|
||||
|
||||
// Test HostInfo
|
||||
info, err := conn.HostInfo()
|
||||
if err != nil {
|
||||
t.Errorf("HostInfo failed: %v", err)
|
||||
}
|
||||
if info.Platform == "" {
|
||||
t.Error("HostInfo.Platform is empty")
|
||||
}
|
||||
|
||||
// Test ListTools (should be empty initially, but shouldn't error)
|
||||
tools, err := conn.ListTools()
|
||||
if err != nil {
|
||||
t.Errorf("ListTools failed: %v", err)
|
||||
}
|
||||
t.Logf("ListTools returned %d tools", len(tools))
|
||||
}
|
||||
|
||||
func TestIntegration_ProcessKeeperReconnect(t *testing.T) {
|
||||
socketPath, cleanup := startTestServer(t)
|
||||
defer cleanup()
|
||||
|
||||
secret := "test-secret-keeper"
|
||||
|
||||
// Create a process manually (simulating what Spawn returns)
|
||||
proc := &Process{
|
||||
SocketPath: socketPath,
|
||||
Secret: secret,
|
||||
}
|
||||
|
||||
// Create keeper without respawn (we're testing reconnect, not respawn)
|
||||
keeper := NewProcessKeeper(context.Background(), proc, nil)
|
||||
|
||||
// First dial
|
||||
conn1, err := keeper.Dial()
|
||||
if err != nil {
|
||||
t.Fatalf("first Dial failed: %v", err)
|
||||
}
|
||||
token := conn1.Token()
|
||||
conn1.Close()
|
||||
|
||||
// Second dial should reconnect with same secret
|
||||
conn2, err := keeper.Dial()
|
||||
if err != nil {
|
||||
t.Fatalf("second Dial failed: %v", err)
|
||||
}
|
||||
defer conn2.Close()
|
||||
|
||||
if conn2.Token() != token {
|
||||
t.Errorf("token changed after reconnect: %q -> %q", token, conn2.Token())
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntegration_ConcurrentConnections(t *testing.T) {
|
||||
socketPath, cleanup := startTestServer(t)
|
||||
defer cleanup()
|
||||
|
||||
secret := "test-secret-concurrent"
|
||||
|
||||
// Establish the secret
|
||||
conn0, err := Dial(socketPath, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("initial Dial failed: %v", err)
|
||||
}
|
||||
expectedToken := conn0.Token()
|
||||
conn0.Close()
|
||||
|
||||
// Open multiple concurrent connections
|
||||
const numConns = 5
|
||||
conns := make([]*Conn, numConns)
|
||||
errors := make([]error, numConns)
|
||||
|
||||
for i := 0; i < numConns; i++ {
|
||||
conns[i], errors[i] = Dial(socketPath, secret)
|
||||
}
|
||||
|
||||
// Check all succeeded with same token
|
||||
for i := 0; i < numConns; i++ {
|
||||
if errors[i] != nil {
|
||||
t.Errorf("connection %d failed: %v", i, errors[i])
|
||||
continue
|
||||
}
|
||||
if conns[i].Token() != expectedToken {
|
||||
t.Errorf("connection %d has different token", i)
|
||||
}
|
||||
conns[i].Close()
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntegration_ProcessKeeperRespawn(t *testing.T) {
|
||||
// This test verifies that when a toolsrv process dies,
|
||||
// the ProcessKeeper respawns it with a new secret.
|
||||
|
||||
cwd, _ := os.Getwd()
|
||||
|
||||
// We need to use the real Spawn which requires the binary in PATH
|
||||
// For now, manually simulate what ProcessKeeper does
|
||||
|
||||
// Start first server
|
||||
socketPath1, cleanup1 := startTestServer(t)
|
||||
|
||||
secret1 := "secret-for-first-server"
|
||||
conn1, err := Dial(socketPath1, secret1)
|
||||
if err != nil {
|
||||
cleanup1()
|
||||
t.Fatalf("first Dial failed: %v", err)
|
||||
}
|
||||
token1 := conn1.Token()
|
||||
conn1.Close()
|
||||
|
||||
// Kill the server
|
||||
cleanup1()
|
||||
|
||||
// Start second server (simulates respawn)
|
||||
socketPath2, cleanup2 := startTestServer(t)
|
||||
defer cleanup2()
|
||||
|
||||
// New server should accept a new secret (any secret, since it's fresh)
|
||||
secret2 := "secret-for-second-server"
|
||||
conn2, err := Dial(socketPath2, secret2)
|
||||
if err != nil {
|
||||
t.Fatalf("second Dial failed: %v", err)
|
||||
}
|
||||
token2 := conn2.Token()
|
||||
conn2.Close()
|
||||
|
||||
// Tokens should be different (different servers)
|
||||
if token1 == token2 {
|
||||
t.Error("tokens should be different after respawn")
|
||||
}
|
||||
|
||||
// Old secret should NOT work on new server
|
||||
conn3, err := Dial(socketPath2, secret1)
|
||||
if err == nil {
|
||||
conn3.Close()
|
||||
t.Error("old secret should not work on new server")
|
||||
}
|
||||
|
||||
// Verify ProcessKeeper handles this correctly
|
||||
proc := &Process{
|
||||
SocketPath: socketPath2,
|
||||
Secret: secret2,
|
||||
}
|
||||
keeper := NewProcessKeeper(context.Background(), proc, nil)
|
||||
|
||||
conn4, err := keeper.Dial()
|
||||
if err != nil {
|
||||
t.Fatalf("keeper Dial failed: %v", err)
|
||||
}
|
||||
if conn4.Token() != token2 {
|
||||
t.Error("keeper should reconnect with same token")
|
||||
}
|
||||
conn4.Close()
|
||||
|
||||
t.Logf("Successfully verified: old_token=%s new_token=%s", token1[:8], token2[:8])
|
||||
_ = cwd // silence unused warning
|
||||
}
|
||||
|
|
@ -2,9 +2,7 @@
|
|||
package toolsrv
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"runtime"
|
||||
|
|
@ -14,36 +12,23 @@ import (
|
|||
)
|
||||
|
||||
// Server9P holds the state for the 9P tool server.
|
||||
// The actual 9P protocol handling is done by the generic 9P server code;
|
||||
// this struct just manages agents, tools, and execution.
|
||||
type Server9P struct {
|
||||
mu sync.RWMutex
|
||||
|
||||
secret []byte // shared secret for auth
|
||||
secret string // set on first auth, verified on subsequent
|
||||
token string // session token returned after auth
|
||||
registry *Registry // tool registry
|
||||
sessID string // session ID for registry scoping
|
||||
yolo bool // skip sandbox
|
||||
|
||||
// Agent registration: token → agent info
|
||||
agents map[string]*AgentInfo
|
||||
|
||||
// Process management
|
||||
fs *FS9P
|
||||
}
|
||||
|
||||
// AgentInfo holds registered agent data.
|
||||
type AgentInfo struct {
|
||||
ID string // immutable agent ID (UUID)
|
||||
CWD string // working directory
|
||||
Token string // auth token
|
||||
}
|
||||
|
||||
// NewServer9P creates a new 9P tool server.
|
||||
func NewServer9P(secret []byte) *Server9P {
|
||||
func NewServer9P() *Server9P {
|
||||
return &Server9P{
|
||||
secret: secret,
|
||||
agents: make(map[string]*AgentInfo),
|
||||
fs: NewFS9P(""), // cwd set per-agent
|
||||
fs: NewFS9P(""), // cwd set per-agent
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -64,47 +49,33 @@ func (s *Server9P) SetYolo(yolo bool) {
|
|||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
// RegisterAgent registers an agent and returns a token.
|
||||
// Verifies the signature using the shared secret.
|
||||
func (s *Server9P) RegisterAgent(id, cwd, sig string) (string, error) {
|
||||
// Verify signature
|
||||
expected := s.computeSig(id, cwd)
|
||||
if !hmac.Equal([]byte(sig), []byte(expected)) {
|
||||
return "", fmt.Errorf("invalid signature")
|
||||
}
|
||||
|
||||
// Generate token
|
||||
tokenBytes := make([]byte, 16)
|
||||
if _, err := rand.Read(tokenBytes); err != nil {
|
||||
return "", err
|
||||
}
|
||||
token := hex.EncodeToString(tokenBytes)
|
||||
|
||||
// Store agent info
|
||||
// Authenticate handles secret verification.
|
||||
// First call sets the secret; subsequent calls must match it.
|
||||
// Returns (token, nil) on success, ("", error) on failure.
|
||||
func (s *Server9P) Authenticate(clientSecret string) (string, error) {
|
||||
s.mu.Lock()
|
||||
s.agents[token] = &AgentInfo{
|
||||
ID: id,
|
||||
CWD: cwd,
|
||||
Token: token,
|
||||
}
|
||||
s.mu.Unlock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
return token, nil
|
||||
if s.secret == "" {
|
||||
// First auth - set the secret
|
||||
s.secret = clientSecret
|
||||
s.token = randomToken()
|
||||
return s.token, nil
|
||||
}
|
||||
|
||||
// Subsequent auth - verify secret
|
||||
if clientSecret != s.secret {
|
||||
return "", fmt.Errorf("authentication failed")
|
||||
}
|
||||
|
||||
return s.token, nil
|
||||
}
|
||||
|
||||
// GetAgent returns agent info for a token.
|
||||
func (s *Server9P) GetAgent(token string) (*AgentInfo, bool) {
|
||||
// Token returns the current session token (empty if not authenticated).
|
||||
func (s *Server9P) Token() string {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
agent, ok := s.agents[token]
|
||||
return agent, ok
|
||||
}
|
||||
|
||||
// computeSig computes HMAC signature for agent registration.
|
||||
func (s *Server9P) computeSig(id, cwd string) string {
|
||||
h := hmac.New(sha256.New, s.secret)
|
||||
h.Write([]byte(id + "|" + cwd))
|
||||
return hex.EncodeToString(h.Sum(nil))
|
||||
return s.token
|
||||
}
|
||||
|
||||
// HostInfo returns platform info as key=value lines.
|
||||
|
|
@ -122,3 +93,9 @@ func (s *Server9P) BuildTree() *fsedsl.Tree {
|
|||
ctx := ToolsrvCtx{Server: s}
|
||||
return fsedsl.BuildTree(ToolsrvSpec(), ctx)
|
||||
}
|
||||
|
||||
func randomToken() string {
|
||||
b := make([]byte, 16)
|
||||
rand.Read(b)
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,81 +4,78 @@ import (
|
|||
"testing"
|
||||
)
|
||||
|
||||
func TestServer9P_RegisterAgent(t *testing.T) {
|
||||
secret := []byte("test-secret-key-1234567890123456")
|
||||
srv := NewServer9P(secret)
|
||||
func TestServer9P_AuthenticateFirst(t *testing.T) {
|
||||
srv := NewServer9P()
|
||||
|
||||
agentID := "agent-uuid-123"
|
||||
cwd := "/home/user/project"
|
||||
sig := ComputeRegistrationSig(secret, agentID, cwd)
|
||||
secret := "test-secret-0123456789abcdef"
|
||||
|
||||
// Register should succeed with valid signature
|
||||
token, err := srv.RegisterAgent(agentID, cwd, sig)
|
||||
// First auth sets the secret and returns a token
|
||||
token, err := srv.Authenticate(secret)
|
||||
if err != nil {
|
||||
t.Fatalf("RegisterAgent failed: %v", err)
|
||||
t.Fatalf("Authenticate failed: %v", err)
|
||||
}
|
||||
if token == "" {
|
||||
t.Fatal("RegisterAgent returned empty token")
|
||||
t.Fatal("Authenticate returned empty token")
|
||||
}
|
||||
|
||||
// Should be able to look up the agent
|
||||
agent, ok := srv.GetAgent(token)
|
||||
if !ok {
|
||||
t.Fatal("GetAgent returned false")
|
||||
}
|
||||
if agent.ID != agentID {
|
||||
t.Errorf("agent.ID = %q, want %q", agent.ID, agentID)
|
||||
}
|
||||
if agent.CWD != cwd {
|
||||
t.Errorf("agent.CWD = %q, want %q", agent.CWD, cwd)
|
||||
// Token() should return the same value
|
||||
if srv.Token() != token {
|
||||
t.Errorf("Token() = %q, want %q", srv.Token(), token)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServer9P_RegisterAgentBadSig(t *testing.T) {
|
||||
secret := []byte("test-secret-key-1234567890123456")
|
||||
srv := NewServer9P(secret)
|
||||
func TestServer9P_AuthenticateSameSecret(t *testing.T) {
|
||||
srv := NewServer9P()
|
||||
|
||||
agentID := "agent-uuid-123"
|
||||
cwd := "/home/user/project"
|
||||
badSig := "invalid-signature"
|
||||
secret := "test-secret-0123456789abcdef"
|
||||
|
||||
// Register should fail with invalid signature
|
||||
_, err := srv.RegisterAgent(agentID, cwd, badSig)
|
||||
// First auth
|
||||
token1, err := srv.Authenticate(secret)
|
||||
if err != nil {
|
||||
t.Fatalf("First Authenticate failed: %v", err)
|
||||
}
|
||||
|
||||
// Second auth with same secret should succeed and return same token
|
||||
token2, err := srv.Authenticate(secret)
|
||||
if err != nil {
|
||||
t.Fatalf("Second Authenticate failed: %v", err)
|
||||
}
|
||||
|
||||
if token1 != token2 {
|
||||
t.Errorf("Token changed: %q -> %q", token1, token2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServer9P_AuthenticateWrongSecret(t *testing.T) {
|
||||
srv := NewServer9P()
|
||||
|
||||
secret := "test-secret-0123456789abcdef"
|
||||
wrongSecret := "wrong-secret-fedcba9876543210"
|
||||
|
||||
// First auth establishes the secret
|
||||
_, err := srv.Authenticate(secret)
|
||||
if err != nil {
|
||||
t.Fatalf("First Authenticate failed: %v", err)
|
||||
}
|
||||
|
||||
// Second auth with wrong secret should fail
|
||||
_, err = srv.Authenticate(wrongSecret)
|
||||
if err == nil {
|
||||
t.Fatal("RegisterAgent should fail with bad signature")
|
||||
t.Fatal("Authenticate should fail with wrong secret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServer9P_RegisterAgentWrongSecret(t *testing.T) {
|
||||
secret := []byte("test-secret-key-1234567890123456")
|
||||
wrongSecret := []byte("wrong-secret-key-abcdefghijklmnop")
|
||||
srv := NewServer9P(secret)
|
||||
func TestServer9P_TokenBeforeAuth(t *testing.T) {
|
||||
srv := NewServer9P()
|
||||
|
||||
agentID := "agent-uuid-123"
|
||||
cwd := "/home/user/project"
|
||||
// Sign with wrong secret
|
||||
sig := ComputeRegistrationSig(wrongSecret, agentID, cwd)
|
||||
|
||||
// Register should fail
|
||||
_, err := srv.RegisterAgent(agentID, cwd, sig)
|
||||
if err == nil {
|
||||
t.Fatal("RegisterAgent should fail with wrong secret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServer9P_GetAgentNotFound(t *testing.T) {
|
||||
secret := []byte("test-secret-key-1234567890123456")
|
||||
srv := NewServer9P(secret)
|
||||
|
||||
_, ok := srv.GetAgent("nonexistent-token")
|
||||
if ok {
|
||||
t.Fatal("GetAgent should return false for nonexistent token")
|
||||
// Token should be empty before authentication
|
||||
if srv.Token() != "" {
|
||||
t.Errorf("Token() = %q before auth, want empty", srv.Token())
|
||||
}
|
||||
}
|
||||
|
||||
func TestServer9P_BuildTree(t *testing.T) {
|
||||
secret := []byte("test-secret-key-1234567890123456")
|
||||
srv := NewServer9P(secret)
|
||||
srv := NewServer9P()
|
||||
|
||||
tree := srv.BuildTree()
|
||||
if tree == nil {
|
||||
|
|
@ -87,11 +84,51 @@ func TestServer9P_BuildTree(t *testing.T) {
|
|||
}
|
||||
|
||||
func TestServer9P_HostInfo(t *testing.T) {
|
||||
secret := []byte("test-secret-key-1234567890123456")
|
||||
srv := NewServer9P(secret)
|
||||
srv := NewServer9P()
|
||||
|
||||
info := srv.HostInfo()
|
||||
if info == "" {
|
||||
t.Fatal("HostInfo returned empty string")
|
||||
}
|
||||
|
||||
// Should contain platform info
|
||||
if !contains(info, "platform=") {
|
||||
t.Error("HostInfo missing platform=")
|
||||
}
|
||||
if !contains(info, "arch=") {
|
||||
t.Error("HostInfo missing arch=")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateSecret(t *testing.T) {
|
||||
s1, err := GenerateSecret()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateSecret failed: %v", err)
|
||||
}
|
||||
if s1 == "" {
|
||||
t.Fatal("GenerateSecret returned empty string")
|
||||
}
|
||||
// 32 bytes -> 64 hex chars
|
||||
if len(s1) != 64 {
|
||||
t.Errorf("GenerateSecret returned %d chars, want 64", len(s1))
|
||||
}
|
||||
|
||||
// Two calls should return different secrets
|
||||
s2, _ := GenerateSecret()
|
||||
if s1 == s2 {
|
||||
t.Error("GenerateSecret returned same secret twice")
|
||||
}
|
||||
}
|
||||
|
||||
func contains(s, substr string) bool {
|
||||
return len(s) >= len(substr) && (s == substr || len(s) > 0 && containsAt(s, substr))
|
||||
}
|
||||
|
||||
func containsAt(s, substr string) bool {
|
||||
for i := 0; i <= len(s)-len(substr); i++ {
|
||||
if s[i:i+len(substr)] == substr {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,243 @@
|
|||
// spawn.go - Process spawning and lifecycle management for toolsrv.
|
||||
package toolsrv
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Process represents a running toolsrv process.
|
||||
type Process struct {
|
||||
Cmd *exec.Cmd
|
||||
SocketPath string
|
||||
Secret string
|
||||
Info ProcessInfo
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// ProcessInfo contains metadata about the toolsrv process.
|
||||
type ProcessInfo struct {
|
||||
Platform string
|
||||
IsGitRepo bool
|
||||
}
|
||||
|
||||
// Option configures spawning behavior.
|
||||
type Option func(*spawnConfig)
|
||||
|
||||
type spawnConfig struct {
|
||||
yolo bool
|
||||
}
|
||||
|
||||
// WithYolo disables sandbox enforcement.
|
||||
func WithYolo() Option {
|
||||
return func(c *spawnConfig) { c.yolo = true }
|
||||
}
|
||||
|
||||
// Spawn starts a local toolsrv process.
|
||||
func Spawn(ctx context.Context, cwd string, opts ...Option) (*Process, error) {
|
||||
cfg := &spawnConfig{}
|
||||
for _, opt := range opts {
|
||||
opt(cfg)
|
||||
}
|
||||
|
||||
// Generate socket path and secret
|
||||
socketPath := filepath.Join(os.TempDir(), fmt.Sprintf("toolsrv-%d-%s.sock", os.Getpid(), randomHex(8)))
|
||||
secret := randomHex(32)
|
||||
|
||||
// Find toolsrv binary
|
||||
toolsrvPath, err := findToolsrv()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("find toolsrv: %w", err)
|
||||
}
|
||||
|
||||
// Build command
|
||||
args := []string{"serve", "--cwd", cwd, "--listen", socketPath}
|
||||
if cfg.yolo {
|
||||
args = append(args, "--yolo")
|
||||
}
|
||||
|
||||
procCtx, cancel := context.WithCancel(ctx)
|
||||
cmd := exec.CommandContext(procCtx, toolsrvPath, args...)
|
||||
// Note: secret is NOT passed via env var. It's established on first
|
||||
// connection via 9P Tauth (socket permissions + SSH handle security).
|
||||
cmd.Stderr = os.Stderr // Let server errors go to stderr for debugging
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
cancel()
|
||||
return nil, fmt.Errorf("start toolsrv: %w", err)
|
||||
}
|
||||
|
||||
// Wait for socket to be ready
|
||||
if err := waitForSocket(socketPath, 5*time.Second); err != nil {
|
||||
cancel()
|
||||
cmd.Process.Kill()
|
||||
return nil, fmt.Errorf("wait for socket: %w", err)
|
||||
}
|
||||
|
||||
return &Process{
|
||||
Cmd: cmd,
|
||||
SocketPath: socketPath,
|
||||
Secret: secret,
|
||||
Info: ProcessInfo{
|
||||
Platform: "linux",
|
||||
IsGitRepo: false, // Caller can override
|
||||
},
|
||||
cancel: cancel,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Kill terminates the process.
|
||||
func (p *Process) Kill() error {
|
||||
if p.cancel != nil {
|
||||
p.cancel()
|
||||
}
|
||||
if p.Cmd != nil && p.Cmd.Process != nil {
|
||||
return p.Cmd.Process.Kill()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close terminates the process (alias for Kill).
|
||||
func (p *Process) Close() error {
|
||||
return p.Kill()
|
||||
}
|
||||
|
||||
// Wait waits for the process to exit.
|
||||
func (p *Process) Wait() error {
|
||||
if p.Cmd != nil {
|
||||
return p.Cmd.Wait()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoteConfig configures remote toolsrv spawning.
|
||||
type RemoteConfig struct {
|
||||
SSHTarget string // user@host or host
|
||||
CWD string
|
||||
}
|
||||
|
||||
// SpawnRemote starts a remote toolsrv process via SSH.
|
||||
func SpawnRemote(ctx context.Context, cfg RemoteConfig) (*Process, error) {
|
||||
// For now, return an error - remote spawning requires SSH bootstrap
|
||||
return nil, fmt.Errorf("remote spawning not yet implemented for 9P")
|
||||
}
|
||||
|
||||
// ProcessKeeper manages process lifecycle with automatic respawning.
|
||||
type ProcessKeeper struct {
|
||||
ctx context.Context
|
||||
proc *Process
|
||||
respawn func(context.Context) (*Process, error)
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// NewProcessKeeper creates a new process keeper.
|
||||
func NewProcessKeeper(ctx context.Context, proc *Process, respawn func(context.Context) (*Process, error)) *ProcessKeeper {
|
||||
return &ProcessKeeper{
|
||||
ctx: ctx,
|
||||
proc: proc,
|
||||
respawn: respawn,
|
||||
}
|
||||
}
|
||||
|
||||
// SetContext updates the context used for respawning.
|
||||
func (pk *ProcessKeeper) SetContext(ctx context.Context) {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
pk.ctx = ctx
|
||||
}
|
||||
|
||||
// Dial connects to the managed process, respawning if necessary.
|
||||
func (pk *ProcessKeeper) Dial() (*Conn, error) {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
|
||||
// Try to connect to existing process
|
||||
if pk.proc != nil {
|
||||
conn, err := Dial(pk.proc.SocketPath, pk.proc.Secret)
|
||||
if err == nil {
|
||||
// Save the secret if this was first connection
|
||||
if pk.proc.Secret == "" {
|
||||
pk.proc.Secret = conn.Secret()
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
// Process may have died, try to respawn
|
||||
}
|
||||
|
||||
if pk.respawn != nil {
|
||||
proc, err := pk.respawn(pk.ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("respawn: %w", err)
|
||||
}
|
||||
pk.proc = proc
|
||||
conn, err := Dial(proc.SocketPath, proc.Secret)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("no process available")
|
||||
}
|
||||
|
||||
// Close terminates the managed process.
|
||||
func (pk *ProcessKeeper) Close() error {
|
||||
pk.mu.Lock()
|
||||
defer pk.mu.Unlock()
|
||||
if pk.proc != nil {
|
||||
return pk.proc.Kill()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// --- Helpers ---
|
||||
|
||||
func randomHex(n int) string {
|
||||
b := make([]byte, n)
|
||||
rand.Read(b)
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
|
||||
func findToolsrv() (string, error) {
|
||||
// Check if toolsrv is in PATH
|
||||
if path, err := exec.LookPath("toolsrv"); err == nil {
|
||||
return path, nil
|
||||
}
|
||||
|
||||
// Check common locations
|
||||
home, _ := os.UserHomeDir()
|
||||
candidates := []string{
|
||||
filepath.Join(home, ".config", "ollie", "bin", "toolsrv"),
|
||||
filepath.Join(home, "go", "bin", "toolsrv"),
|
||||
"/usr/local/bin/toolsrv",
|
||||
}
|
||||
|
||||
for _, p := range candidates {
|
||||
if _, err := os.Stat(p); err == nil {
|
||||
return p, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("toolsrv binary not found in PATH or common locations")
|
||||
}
|
||||
|
||||
func waitForSocket(path string, timeout time.Duration) error {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
conn, err := net.DialTimeout("unix", path, 100*time.Millisecond)
|
||||
if err == nil {
|
||||
conn.Close()
|
||||
return nil
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
return fmt.Errorf("timeout waiting for socket %s", path)
|
||||
}
|
||||
|
|
@ -14,8 +14,7 @@ import (
|
|||
// ToolsrvCtx is the context passed through the fsedsl tree.
|
||||
type ToolsrvCtx struct {
|
||||
Server *Server9P
|
||||
Agent *AgentInfo // set when traversing into agent-specific paths
|
||||
Proc *Proc9P // set when traversing into /proc/{pid}
|
||||
Proc *Proc9P // set when traversing into /proc/{pid}
|
||||
}
|
||||
|
||||
// Type aliases from fsedsl, specialized to ToolsrvCtx.
|
||||
|
|
@ -37,12 +36,9 @@ var (
|
|||
)
|
||||
|
||||
// ToolsrvSpec returns the fsedsl specification for the toolsrv namespace.
|
||||
// Authentication is handled via 9P Tauth, not a /register file.
|
||||
func ToolsrvSpec() FsNodeDecl {
|
||||
return DirNode("/",
|
||||
FileNode("register", 0666,
|
||||
Doc("Register agent: write aname=<id>\\ncwd=<path>\\nsig=<sig>, read token"),
|
||||
Request(handleRegister),
|
||||
),
|
||||
FileNode("ctl", 0222,
|
||||
Doc("Control: write 'load <tool>' or 'unload <tool>'"),
|
||||
Write(handleCtlWrite),
|
||||
|
|
@ -89,23 +85,6 @@ func ToolsrvSpec() FsNodeDecl {
|
|||
|
||||
// --- Handlers ---
|
||||
|
||||
func handleRegister(ctx ToolsrvCtx, data []byte) ([]byte, error) {
|
||||
args := parseKV(string(data))
|
||||
id := args["aname"]
|
||||
cwd := args["cwd"]
|
||||
sig := args["sig"]
|
||||
|
||||
if id == "" || cwd == "" || sig == "" {
|
||||
return nil, fmt.Errorf("missing aname, cwd, or sig")
|
||||
}
|
||||
|
||||
token, err := ctx.Server.RegisterAgent(id, cwd, sig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []byte(token + "\n"), nil
|
||||
}
|
||||
|
||||
func handleCtlWrite(ctx ToolsrvCtx, data []byte) error {
|
||||
return ctx.Server.fs.HandleCtl(strings.TrimSpace(string(data)))
|
||||
}
|
||||
|
|
@ -124,22 +103,13 @@ func handleInfoRead(ctx ToolsrvCtx) ([]byte, error) {
|
|||
|
||||
func handleProcNew(ctx ToolsrvCtx, data []byte) ([]byte, error) {
|
||||
args := parseKV(string(data))
|
||||
token := args["token"]
|
||||
|
||||
agent, ok := ctx.Server.GetAgent(token)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("invalid token")
|
||||
}
|
||||
// Token verification is now done at connection level via Tauth.
|
||||
// Just pass through to the FS handler.
|
||||
delete(args, "token") // remove token from payload if present
|
||||
|
||||
// Set cwd for this execution
|
||||
ctx.Server.fs.SetCWD(agent.CWD)
|
||||
|
||||
// Build payload without token
|
||||
var payload strings.Builder
|
||||
for k, v := range args {
|
||||
if k == "token" {
|
||||
continue
|
||||
}
|
||||
payload.WriteString(k)
|
||||
payload.WriteString("=")
|
||||
payload.WriteString(v)
|
||||
|
|
@ -155,22 +125,12 @@ func handleProcNew(ctx ToolsrvCtx, data []byte) ([]byte, error) {
|
|||
|
||||
func handleProcNewBg(ctx ToolsrvCtx, data []byte) error {
|
||||
args := parseKV(string(data))
|
||||
token := args["token"]
|
||||
|
||||
agent, ok := ctx.Server.GetAgent(token)
|
||||
if !ok {
|
||||
return fmt.Errorf("invalid token")
|
||||
}
|
||||
// Token verification is now done at connection level via Tauth.
|
||||
delete(args, "token") // remove token from payload if present
|
||||
|
||||
// Set cwd for this execution
|
||||
ctx.Server.fs.SetCWD(agent.CWD)
|
||||
|
||||
// Build payload without token
|
||||
var payload strings.Builder
|
||||
for k, v := range args {
|
||||
if k == "token" {
|
||||
continue
|
||||
}
|
||||
payload.WriteString(k)
|
||||
payload.WriteString("=")
|
||||
payload.WriteString(v)
|
||||
|
|
|
|||
Loading…
Reference in New Issue