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:
parent
659f89fcbf
commit
de5456e4aa
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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 {
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue