ollie/cmd/olliesrv/internal/toolclient/toolsrv.go

622 lines
15 KiB
Go

// client9p.go - 9P client for toolsrvclient.
package toolclient
import (
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"strings"
"syscall"
"9fans.net/go/plan9"
p9client "9fans.net/go/plan9/client"
"ollie/toolsrv/protocol"
)
// ToolsrvConn is a 9P connection to a toolsrvclient.
type ToolsrvConn struct {
fsys *p9client.Fsys
conn *p9client.Conn
token string // session token from auth
secret string // secret used for auth (for reconnect)
agentID string // agent identity for per-agent tool registry
onToolsChanged func() // callback for tool changes (client must poll)
}
// SetAgentID sets the agent identity used for tool registry scoping.
func (c *ToolsrvConn) SetAgentID(id string) {
c.agentID = id
}
// 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 DialToolsrv(socketPath string, secret string) (*ToolsrvConn, error) {
conn, err := p9client.Dial("unix", socketPath)
if err != nil {
return nil, err
}
// Generate secret if not provided
if secret == "" {
secret, err = randomSecret()
if err != nil {
conn.Close()
return nil, fmt.Errorf("generate secret: %w", err)
}
}
// 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 &ToolsrvConn{
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 *ToolsrvConn) Secret() string {
return c.secret
}
// Token returns the session token.
func (c *ToolsrvConn) Token() string {
return c.token
}
// ListTools returns the list of loaded tools for this agent.
func (c *ToolsrvConn) ListTools() ([]protocol.ToolInfo, error) {
fid, err := c.fsys.Open("tools", plan9.ORDWR)
if err != nil {
return nil, err
}
defer fid.Close()
if _, err := fid.Write([]byte(c.agentID)); err != nil {
return nil, fmt.Errorf("write agent id: %w", err)
}
var result []byte
buf := make([]byte, 8192)
for offset := int64(0); ; {
n, err := fid.ReadAt(buf, offset)
if n > 0 {
result = append(result, buf[:n]...)
offset += int64(n)
}
if err == io.EOF || n == 0 {
break
}
if err != nil {
return nil, fmt.Errorf("read tools: %w", err)
}
}
var tools []protocol.ToolInfo
if err := json.Unmarshal(result, &tools); err != nil {
return nil, fmt.Errorf("parse tools: %w", err)
}
return tools, nil
}
// CallTool executes a tool and returns the result.
func (c *ToolsrvConn) 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)
}
// Close fid on context cancellation to unblock the read
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-ctx.Done():
fid.Close()
case <-done:
}
}()
defer fid.Close()
// Convert JSON args to key=value format
var argMap map[string]any
if err := json.Unmarshal(args, &argMap); err != nil {
argMap = make(map[string]any)
}
var payload strings.Builder
fmt.Fprintf(&payload, "token=%s\ntool=%s\nagent=%s\n", c.token, name, c.agentID)
for k, v := range argMap {
escaped := escapeValue(fmt.Sprintf("%v", v))
fmt.Fprintf(&payload, "%s=%s\n", k, escaped)
}
if _, err := fid.Write([]byte(payload.String())); err != nil {
if ctx.Err() != nil {
return nil, ctx.Err()
}
return nil, fmt.Errorf("write: %w", err)
}
// Read result from offset 0 (Write advances the file offset, but
// for request-response files we need to read the result from the start)
var result []byte
buf := make([]byte, 8192)
for offset := int64(0); ; {
n, err := fid.ReadAt(buf, offset)
if n > 0 {
result = append(result, buf[:n]...)
offset += int64(n)
}
if err == io.EOF || n == 0 {
break
}
if err != nil {
if ctx.Err() != nil {
return nil, ctx.Err()
}
return nil, fmt.Errorf("read: %w", err)
}
}
return json.RawMessage(result), nil
}
// CallToolBackground executes a tool in the background, returning the PID immediately.
// Use GetDetachedOutput(pid) to read output later.
func (c *ToolsrvConn) CallToolBackground(name string, args json.RawMessage) (int, error) {
fid, err := c.fsys.Open("proc/new.bg", plan9.ORDWR)
if err != nil {
return 0, fmt.Errorf("open proc/new.bg: %w", err)
}
defer fid.Close()
var argMap map[string]any
if err := json.Unmarshal(args, &argMap); err != nil {
argMap = make(map[string]any)
}
var payload strings.Builder
fmt.Fprintf(&payload, "token=%s\ntool=%s\nagent=%s\n", c.token, name, c.agentID)
for k, v := range argMap {
escaped := escapeValue(fmt.Sprintf("%v", v))
fmt.Fprintf(&payload, "%s=%s\n", k, escaped)
}
if _, err := fid.Write([]byte(payload.String())); err != nil {
return 0, fmt.Errorf("write: %w", err)
}
// Read PID from response
buf := make([]byte, 64)
n, err := fid.ReadAt(buf, 0)
if err != nil && err != io.EOF {
return 0, fmt.Errorf("read pid: %w", err)
}
var pid int
fmt.Sscanf(string(buf[:n]), "%d", &pid)
return pid, nil
}
// LoadTool loads a tool by name.
func (c *ToolsrvConn) 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 " + c.agentID + " " + name + "\n"))
return err
}
// ListAllTools reads the full catalog of available tools from toolsrvclient.
func (c *ToolsrvConn) ListAllTools() ([]byte, error) {
fid, err := c.fsys.Open("all", plan9.OREAD)
if err != nil {
return nil, err
}
defer fid.Close()
return io.ReadAll(fid)
}
// UnloadTool unloads a tool by name.
func (c *ToolsrvConn) 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 " + c.agentID + " " + name + "\n"))
return err
}
// ClearTools removes all tools for this agent.
func (c *ToolsrvConn) ClearTools() error {
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
if err != nil {
return err
}
defer fid.Close()
_, err = fid.Write([]byte("clear " + c.agentID + "\n"))
return err
}
// Ping checks if the connection is alive.
func (c *ToolsrvConn) Ping() error {
_, err := c.fsys.Stat("info")
return err
}
// Close closes the connection.
func (c *ToolsrvConn) Close() {
if c.conn != nil {
c.conn.Close()
}
}
// HostInfo returns host information from the server.
func (c *ToolsrvConn) 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 strings.Split(string(data), "\n") {
if idx := strings.IndexByte(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 *ToolsrvConn) SetEnv(key, value string) error {
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
if err != nil {
return err
}
defer fid.Close()
_, err = fmt.Fprintf(fid, "env %s=%s\n", key, value)
return err
}
// SetCWD sets the working directory via ctl command.
func (c *ToolsrvConn) SetCWD(dir string) error {
fid, err := c.fsys.Open("ctl", plan9.OWRITE)
if err != nil {
return err
}
defer fid.Close()
_, err = fmt.Fprintf(fid, "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 *ToolsrvConn) SetOnToolsChanged(fn func()) {
c.onToolsChanged = fn
}
// ToolRegistryRevision returns the current tool registry revision.
func (c *ToolsrvConn) ToolRegistryRevision() uint64 {
fid, err := c.fsys.Open("tools_rev", plan9.ORDWR)
if err != nil {
return 0
}
defer fid.Close()
if _, err := fid.Write([]byte(c.agentID)); err != nil {
return 0
}
data, err := io.ReadAll(fid)
if err != nil {
return 0
}
var rev uint64
fmt.Sscanf(string(data), "%d", &rev)
return rev
}
// --- Process management ---
// Detach detaches the current foreground process to background.
// Returns true if successful.
func (c *ToolsrvConn) 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
}
// ListProcs returns the proc listing for this agent (filtered by agent ID).
func (c *ToolsrvConn) ListProcs() ([]byte, error) {
fid, err := c.fsys.Open("proc/list", plan9.ORDWR)
if err != nil {
return nil, err
}
defer fid.Close()
if _, err := fid.Write([]byte(c.agentID)); err != nil {
return nil, err
}
if _, err := fid.Seek(0, 0); err != nil {
return nil, err
}
return io.ReadAll(fid)
}
// ListProcsIdx returns a machine-readable TSV proc listing.
// Format: pid<TAB>status<TAB>exit_code<TAB>tool<TAB>cmd<TAB>start_unix<TAB>runtime_sec
func (c *ToolsrvConn) ListProcsIdx() ([]byte, error) {
fid, err := c.fsys.Open("proc/idx", plan9.ORDWR)
if err != nil {
return nil, err
}
defer fid.Close()
if _, err := fid.Write([]byte(c.agentID)); err != nil {
return nil, err
}
if _, err := fid.Seek(0, 0); err != nil {
return nil, err
}
return io.ReadAll(fid)
}
// readProcInfo reads info about a single proc.
func (c *ToolsrvConn) 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 strings.Split(string(data), "\n") {
if idx := strings.IndexByte(line, '='); idx >= 0 {
key := line[:idx]
val := line[idx+1:]
switch key {
case "id":
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 *ToolsrvConn) 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 = fmt.Fprintf(fid, "signal %d\n", sig)
return err
}
// DismissDetached dismisses a detached process.
func (c *ToolsrvConn) 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 *ToolsrvConn) 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 toolsrvclient.
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()
}
// --- Helpers ---
func randomSecret() (string, error) {
b := make([]byte, 16)
if _, err := rand.Read(b); err != nil {
return "", err
}
return hex.EncodeToString(b), nil
}
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)
}
// ReadBypassPending blocks until a bypass request is available.
// This should be called in a loop by the approval handler.
// Returns nil, ctx.Err() if context is cancelled.
func (c *ToolsrvConn) ReadBypassPending(ctx context.Context) (*protocol.BypassRequest, error) {
fid, err := c.fsys.Open("bypass/pending", plan9.OREAD)
if err != nil {
return nil, fmt.Errorf("open bypass/pending: %w", err)
}
// Watch for context cancellation to unblock the read
done := make(chan struct{})
defer close(done)
go func() {
select {
case <-ctx.Done():
fid.Close()
case <-done:
}
}()
// Blocking read - will return when a request is available or fid is closed
var result []byte
buf := make([]byte, 8192)
for {
n, err := fid.Read(buf)
if n > 0 {
result = append(result, buf[:n]...)
}
if err == io.EOF || n == 0 {
break
}
if err != nil {
fid.Close()
if ctx.Err() != nil {
return nil, ctx.Err()
}
return nil, fmt.Errorf("read bypass/pending: %w", err)
}
}
fid.Close()
var req protocol.BypassRequest
if err := json.Unmarshal(result, &req); err != nil {
return nil, fmt.Errorf("parse bypass request: %w", err)
}
return &req, nil
}
// ResolveBypass sends the approval/denial decision for a bypass request.
func (c *ToolsrvConn) ResolveBypass(id string, approved bool, errMsg string) error {
fid, err := c.fsys.Open("bypass/resolve", plan9.OWRITE)
if err != nil {
return fmt.Errorf("open bypass/resolve: %w", err)
}
defer fid.Close()
msg := struct {
ID string `json:"id"`
Approved bool `json:"approved"`
Error string `json:"error,omitempty"`
}{ID: id, Approved: approved, Error: errMsg}
data, err := json.Marshal(msg)
if err != nil {
return fmt.Errorf("marshal resolve: %w", err)
}
if _, err := fid.Write(data); err != nil {
return fmt.Errorf("write bypass/resolve: %w", err)
}
return nil
}