toolsrv: fix per-agent tool registry to actually work

The registry was keyed by agent ID but read it from a shared env map
(st.env["OLLIE_UNAME"]) that all agents overwrote — making it
effectively per-session with a race condition.

Fix: pass agent ID explicitly through the protocol at every call site.

- ctl protocol: 'load <agentID> <tool>', 'unload <agentID> <tool>'
- tools file: rdwr (Request) — write agent ID, read filtered list
- proc/new payload: 'agent=<id>' field required
- Client: SetAgentID() stores identity, included in all operations
- Empty agent ID is a hard error everywhere
- Setting OLLIE_UNAME via env ctl is blocked (prevents reintroduction)
This commit is contained in:
Ollie Agent 2026-08-11 16:26:12 +02:00
parent 659f89fcbf
commit de5456e4aa
7 changed files with 95 additions and 60 deletions

View File

@ -393,7 +393,7 @@ func (ag *Agent) SetSessionEnv(sessionID string) {
} }
ag.runtime.ToolServer.SetEnv("OLLIE_SESSION_ID", sessionID) ag.runtime.ToolServer.SetEnv("OLLIE_SESSION_ID", sessionID)
if ag.id != "" { if ag.id != "" {
ag.runtime.ToolServer.SetEnv("OLLIE_UNAME", ag.id) ag.runtime.ToolServer.SetAgentID(ag.id)
} }
if u := os.Getenv("USER"); u != "" { if u := os.Getenv("USER"); u != "" {
ag.runtime.ToolServer.SetEnv("USER", u) ag.runtime.ToolServer.SetEnv("USER", u)

View File

@ -184,7 +184,7 @@ func LoadAutoLoadTools(cfg *agent.AgentConfig, conn *toolsrv.Conn, sessID, uname
conn.SetEnv("OLLIE_SESSION_ID", sessID) conn.SetEnv("OLLIE_SESSION_ID", sessID)
} }
if uname != "" { if uname != "" {
conn.SetEnv("OLLIE_UNAME", uname) conn.SetAgentID(uname)
} }
for _, tl := range cfg.AutoLoad { for _, tl := range cfg.AutoLoad {

View File

@ -191,6 +191,7 @@ func TestIntegration_BasicOperations(t *testing.T) {
t.Fatalf("Dial failed: %v", err) t.Fatalf("Dial failed: %v", err)
} }
defer conn.Close() defer conn.Close()
conn.SetAgentID("test-agent")
// Test Ping // Test Ping
if err := conn.Ping(); err != nil { if err := conn.Ping(); err != nil {
@ -228,6 +229,7 @@ func TestIntegration_ToolExecution(t *testing.T) {
t.Fatalf("Dial failed: %v", err) t.Fatalf("Dial failed: %v", err)
} }
defer conn.Close() defer conn.Close()
conn.SetAgentID("test-agent")
// Load the shell tool // Load the shell tool
if err := conn.LoadTool("shell"); err != nil { if err := conn.LoadTool("shell"); err != nil {
@ -295,6 +297,7 @@ func TestIntegration_MultilineContent(t *testing.T) {
t.Fatalf("Dial failed: %v", err) t.Fatalf("Dial failed: %v", err)
} }
defer conn.Close() defer conn.Close()
conn.SetAgentID("test-agent")
// Load the file_write tool // Load the file_write tool
if err := conn.LoadTool("file_write"); err != nil { if err := conn.LoadTool("file_write"); err != nil {
@ -348,6 +351,7 @@ func TestIntegration_ToolExecutionError(t *testing.T) {
t.Fatalf("Dial failed: %v", err) t.Fatalf("Dial failed: %v", err)
} }
defer conn.Close() defer conn.Close()
conn.SetAgentID("test-agent")
// Load the shell tool // Load the shell tool
if err := conn.LoadTool("shell"); err != nil { if err := conn.LoadTool("shell"); err != nil {

View File

@ -117,36 +117,35 @@ func (st *State) GetEnv(key string) string {
// --- Tool Registry --- // --- Tool Registry ---
// agentID returns the current agent identity from the environment. // ListTools returns loaded tools for the given agent.
func (st *State) agentID() string { func (st *State) ListTools(agentID string) []toolsrv.ToolInfo {
return st.env["OLLIE_UNAME"] if agentID == "" {
} return nil
}
// ListTools returns loaded tools for the current agent.
func (st *State) ListTools() []toolsrv.ToolInfo {
st.mu.RLock() st.mu.RLock()
reg := st.registry reg := st.registry
aid := st.agentID()
st.mu.RUnlock() st.mu.RUnlock()
if reg == nil { if reg == nil {
return nil return nil
} }
return reg.Loaded(aid) return reg.Loaded(agentID)
} }
// LoadTool loads a tool by name for the current agent. // LoadTool loads a tool by name for the given agent.
func (st *State) LoadTool(name string) error { func (st *State) LoadTool(agentID, name string) error {
if agentID == "" {
return fmt.Errorf("agent ID required")
}
st.mu.RLock() st.mu.RLock()
reg := st.registry reg := st.registry
aid := st.agentID()
st.mu.RUnlock() st.mu.RUnlock()
if reg == nil { if reg == nil {
return fmt.Errorf("no tool registry configured") return fmt.Errorf("no tool registry configured")
} }
if err := reg.Load(aid, name); err != nil { if err := reg.Load(agentID, name); err != nil {
return err return err
} }
if st.OnToolsChanged != nil { if st.OnToolsChanged != nil {
@ -155,18 +154,20 @@ func (st *State) LoadTool(name string) error {
return nil return nil
} }
// UnloadTool removes a tool from the current agent's registry. // UnloadTool removes a tool from the given agent's registry.
func (st *State) UnloadTool(name string) error { func (st *State) UnloadTool(agentID, name string) error {
if agentID == "" {
return fmt.Errorf("agent ID required")
}
st.mu.RLock() st.mu.RLock()
reg := st.registry reg := st.registry
aid := st.agentID()
st.mu.RUnlock() st.mu.RUnlock()
if reg == nil { if reg == nil {
return fmt.Errorf("no tool registry configured") return fmt.Errorf("no tool registry configured")
} }
if err := reg.Unload(aid, name); err != nil { if err := reg.Unload(agentID, name); err != nil {
return err return err
} }
if st.OnToolsChanged != nil { if st.OnToolsChanged != nil {
@ -208,7 +209,6 @@ func (st *State) NewProc(ctx context.Context, payload string, background bool) (
// Look up tool // Look up tool
st.mu.RLock() st.mu.RLock()
reg := st.registry reg := st.registry
aid := st.agentID()
cwd := st.cwd cwd := st.cwd
yolo := st.yolo yolo := st.yolo
envCopy := make(map[string]string) envCopy := make(map[string]string)
@ -217,6 +217,12 @@ func (st *State) NewProc(ctx context.Context, payload string, background bool) (
} }
st.mu.RUnlock() st.mu.RUnlock()
aid := args["agent"]
if aid == "" {
return "", 0, fmt.Errorf("missing 'agent' in payload")
}
envCopy["OLLIE_UNAME"] = aid
if reg == nil { if reg == nil {
return "", 0, fmt.Errorf("tool registry not configured") return "", 0, fmt.Errorf("tool registry not configured")
} }
@ -456,15 +462,15 @@ func (st *State) HandleCtl(input string) error {
switch parts[0] { switch parts[0] {
case "load": case "load":
if len(parts) < 2 { if len(parts) < 3 {
return fmt.Errorf("load requires tool name") return fmt.Errorf("load requires agent ID and tool name")
} }
return st.LoadTool(parts[1]) return st.LoadTool(parts[1], parts[2])
case "unload": case "unload":
if len(parts) < 2 { if len(parts) < 3 {
return fmt.Errorf("unload requires tool name") return fmt.Errorf("unload requires agent ID and tool name")
} }
return st.UnloadTool(parts[1]) return st.UnloadTool(parts[1], parts[2])
case "env": case "env":
if len(parts) < 2 { if len(parts) < 2 {
return fmt.Errorf("env requires KEY=VALUE") return fmt.Errorf("env requires KEY=VALUE")
@ -475,7 +481,11 @@ func (st *State) HandleCtl(input string) error {
if idx < 0 { if idx < 0 {
return fmt.Errorf("env requires KEY=VALUE format") return fmt.Errorf("env requires KEY=VALUE format")
} }
st.SetEnv(kv[:idx], kv[idx+1:]) key := kv[:idx]
if key == "OLLIE_UNAME" {
return fmt.Errorf("OLLIE_UNAME is set per-request via agent ID, not via env")
}
st.SetEnv(key, kv[idx+1:])
return nil return nil
case "cwd": case "cwd":
if len(parts) < 2 { if len(parts) < 2 {
@ -488,18 +498,13 @@ func (st *State) HandleCtl(input string) error {
} }
} }
// HandleToolsRead returns the tool list as JSON. // HandleToolsRequest handles the tools rdwr: write agent ID, read tool list.
func (st *State) HandleToolsRead() string { func (st *State) HandleToolsRequest(agentID string) string {
tools := st.ListTools() tools := st.ListTools(agentID)
data, _ := json.Marshal(tools) data, _ := json.Marshal(tools)
return string(data) return string(data)
} }
// HandleToolsWrite loads a tool by name.
func (st *State) HandleToolsWrite(name string) error {
return st.LoadTool(strings.TrimSpace(name))
}
// HandleProcNew is the rdwr handler for /proc/new (blocking). // HandleProcNew is the rdwr handler for /proc/new (blocking).
func (st *State) HandleProcNew(ctx context.Context, payload string) (string, error) { func (st *State) HandleProcNew(ctx context.Context, payload string) (string, error) {
result, _, err := st.NewProc(ctx, payload, false) result, _, err := st.NewProc(ctx, payload, false)

View File

@ -37,13 +37,13 @@ func TestState_ToolsWithoutRegistry(t *testing.T) {
st := NewState("/tmp") st := NewState("/tmp")
// Without registry, ListTools should return nil // Without registry, ListTools should return nil
tools := st.ListTools() tools := st.ListTools("agent1")
if tools != nil { if tools != nil {
t.Errorf("ListTools() without registry = %v, want nil", tools) t.Errorf("ListTools() without registry = %v, want nil", tools)
} }
// Load should fail without registry // Load should fail without registry
err := st.LoadTool("shell") err := st.LoadTool("agent1", "shell")
if err == nil { if err == nil {
t.Error("LoadTool without registry should fail") t.Error("LoadTool without registry should fail")
} }
@ -58,16 +58,22 @@ func TestState_HandleCtl(t *testing.T) {
t.Error("HandleCtl(unknown) should fail") t.Error("HandleCtl(unknown) should fail")
} }
// load without name // load without agent and name
err = st.HandleCtl("load") err = st.HandleCtl("load")
if err == nil { if err == nil {
t.Error("HandleCtl(load) without name should fail") t.Error("HandleCtl(load) without args should fail")
} }
// unload without name // load with only one arg (missing tool name)
err = st.HandleCtl("load agent1")
if err == nil {
t.Error("HandleCtl(load agent1) without tool name should fail")
}
// unload without agent and name
err = st.HandleCtl("unload") err = st.HandleCtl("unload")
if err == nil { if err == nil {
t.Error("HandleCtl(unload) without name should fail") t.Error("HandleCtl(unload) without args should fail")
} }
} }

View File

@ -43,9 +43,8 @@ func Spec() FsNodeDecl {
Write(handleCtlWrite), Write(handleCtlWrite),
), ),
FileNode("tools", 0666, FileNode("tools", 0666,
Doc("Tools: read to list, write tool name to load"), Doc("Tools: write agent ID, read tool list (rdwr)"),
Read(handleToolsRead), Request(handleToolsRequest),
Write(handleToolsWrite),
), ),
FileNode("info", 0444, FileNode("info", 0444,
Doc("Host info: platform, arch"), Doc("Host info: platform, arch"),
@ -88,12 +87,12 @@ func handleCtlWrite(ctx Ctx, data []byte) error {
return ctx.Server.Fs.HandleCtl(strings.TrimSpace(string(data))) return ctx.Server.Fs.HandleCtl(strings.TrimSpace(string(data)))
} }
func handleToolsRead(ctx Ctx) ([]byte, error) { func handleToolsRequest(ctx Ctx, data []byte) ([]byte, error) {
return []byte(ctx.Server.Fs.HandleToolsRead()), nil agentID := strings.TrimSpace(string(data))
} if agentID == "" {
return nil, fmt.Errorf("agent ID required")
func handleToolsWrite(ctx Ctx, data []byte) error { }
return ctx.Server.Fs.HandleToolsWrite(strings.TrimSpace(string(data))) return []byte(ctx.Server.Fs.HandleToolsRequest(agentID)), nil
} }
func handleInfoRead(ctx Ctx) ([]byte, error) { func handleInfoRead(ctx Ctx) ([]byte, error) {

View File

@ -20,9 +20,15 @@ type Conn struct {
conn *p9client.Conn conn *p9client.Conn
token string // session token from auth token string // session token from auth
secret string // secret used for auth (for reconnect) 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) onToolsChanged func() // callback for tool changes (client must poll)
} }
// SetAgentID sets the agent identity used for tool registry scoping.
func (c *Conn) SetAgentID(id string) {
c.agentID = id
}
// Dial connects to a toolsrv at the given Unix socket path and authenticates. // Dial connects to a toolsrv at the given Unix socket path and authenticates.
// If secret is empty, a random one is generated (first connection). // If secret is empty, a random one is generated (first connection).
// Returns the connection and the secret used (caller should save for reconnect). // Returns the connection and the secret used (caller should save for reconnect).
@ -95,21 +101,36 @@ func (c *Conn) Token() string {
return c.token return c.token
} }
// ListTools returns the list of loaded tools. // ListTools returns the list of loaded tools for this agent.
func (c *Conn) ListTools() ([]ToolInfo, error) { func (c *Conn) ListTools() ([]ToolInfo, error) {
fid, err := c.fsys.Open("tools", plan9.OREAD) fid, err := c.fsys.Open("tools", plan9.ORDWR)
if err != nil { if err != nil {
return nil, err return nil, err
} }
defer fid.Close() defer fid.Close()
data, err := io.ReadAll(fid) if _, err := fid.Write([]byte(c.agentID)); err != nil {
if err != nil { return nil, fmt.Errorf("write agent id: %w", err)
return nil, 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 []ToolInfo var tools []ToolInfo
if err := json.Unmarshal(data, &tools); err != nil { if err := json.Unmarshal(result, &tools); err != nil {
return nil, fmt.Errorf("parse tools: %w", err) return nil, fmt.Errorf("parse tools: %w", err)
} }
return tools, nil return tools, nil
@ -129,7 +150,7 @@ func (c *Conn) CallTool(ctx context.Context, name string, args json.RawMessage)
argMap = make(map[string]interface{}) argMap = make(map[string]interface{})
} }
payload := fmt.Sprintf("token=%s\ntool=%s\n", c.token, name) payload := fmt.Sprintf("token=%s\ntool=%s\nagent=%s\n", c.token, name, c.agentID)
for k, v := range argMap { for k, v := range argMap {
escaped := escapeValue(fmt.Sprintf("%v", v)) escaped := escapeValue(fmt.Sprintf("%v", v))
payload += fmt.Sprintf("%s=%s\n", k, escaped) payload += fmt.Sprintf("%s=%s\n", k, escaped)
@ -174,7 +195,7 @@ func (c *Conn) CallToolBackground(name string, args json.RawMessage) (int, error
argMap = make(map[string]interface{}) argMap = make(map[string]interface{})
} }
payload := fmt.Sprintf("token=%s\ntool=%s\n", c.token, name) payload := fmt.Sprintf("token=%s\ntool=%s\nagent=%s\n", c.token, name, c.agentID)
for k, v := range argMap { for k, v := range argMap {
escaped := escapeValue(fmt.Sprintf("%v", v)) escaped := escapeValue(fmt.Sprintf("%v", v))
payload += fmt.Sprintf("%s=%s\n", k, escaped) payload += fmt.Sprintf("%s=%s\n", k, escaped)
@ -203,7 +224,7 @@ func (c *Conn) LoadTool(name string) error {
return err return err
} }
defer fid.Close() defer fid.Close()
_, err = fid.Write([]byte("load " + name + "\n")) _, err = fid.Write([]byte("load " + c.agentID + " " + name + "\n"))
return err return err
} }
@ -214,7 +235,7 @@ func (c *Conn) UnloadTool(name string) error {
return err return err
} }
defer fid.Close() defer fid.Close()
_, err = fid.Write([]byte("unload " + name + "\n")) _, err = fid.Write([]byte("unload " + c.agentID + " " + name + "\n"))
return err return err
} }