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)
|
||||
if ag.id != "" {
|
||||
ag.runtime.ToolServer.SetEnv("OLLIE_UNAME", ag.id)
|
||||
ag.runtime.ToolServer.SetAgentID(ag.id)
|
||||
}
|
||||
if u := os.Getenv("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)
|
||||
}
|
||||
if uname != "" {
|
||||
conn.SetEnv("OLLIE_UNAME", uname)
|
||||
conn.SetAgentID(uname)
|
||||
}
|
||||
|
||||
for _, tl := range cfg.AutoLoad {
|
||||
|
|
|
|||
|
|
@ -191,6 +191,7 @@ func TestIntegration_BasicOperations(t *testing.T) {
|
|||
t.Fatalf("Dial failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.SetAgentID("test-agent")
|
||||
|
||||
// Test Ping
|
||||
if err := conn.Ping(); err != nil {
|
||||
|
|
@ -228,6 +229,7 @@ func TestIntegration_ToolExecution(t *testing.T) {
|
|||
t.Fatalf("Dial failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.SetAgentID("test-agent")
|
||||
|
||||
// Load the shell tool
|
||||
if err := conn.LoadTool("shell"); err != nil {
|
||||
|
|
@ -295,6 +297,7 @@ func TestIntegration_MultilineContent(t *testing.T) {
|
|||
t.Fatalf("Dial failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.SetAgentID("test-agent")
|
||||
|
||||
// Load the file_write tool
|
||||
if err := conn.LoadTool("file_write"); err != nil {
|
||||
|
|
@ -348,6 +351,7 @@ func TestIntegration_ToolExecutionError(t *testing.T) {
|
|||
t.Fatalf("Dial failed: %v", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.SetAgentID("test-agent")
|
||||
|
||||
// Load the shell tool
|
||||
if err := conn.LoadTool("shell"); err != nil {
|
||||
|
|
|
|||
|
|
@ -117,36 +117,35 @@ func (st *State) GetEnv(key string) string {
|
|||
|
||||
// --- Tool Registry ---
|
||||
|
||||
// agentID returns the current agent identity from the environment.
|
||||
func (st *State) agentID() string {
|
||||
return st.env["OLLIE_UNAME"]
|
||||
}
|
||||
|
||||
// ListTools returns loaded tools for the current agent.
|
||||
func (st *State) ListTools() []toolsrv.ToolInfo {
|
||||
// ListTools returns loaded tools for the given agent.
|
||||
func (st *State) ListTools(agentID string) []toolsrv.ToolInfo {
|
||||
if agentID == "" {
|
||||
return nil
|
||||
}
|
||||
st.mu.RLock()
|
||||
reg := st.registry
|
||||
aid := st.agentID()
|
||||
st.mu.RUnlock()
|
||||
|
||||
if reg == nil {
|
||||
return nil
|
||||
}
|
||||
return reg.Loaded(aid)
|
||||
return reg.Loaded(agentID)
|
||||
}
|
||||
|
||||
// LoadTool loads a tool by name for the current agent.
|
||||
func (st *State) LoadTool(name string) error {
|
||||
// LoadTool loads a tool by name for the given agent.
|
||||
func (st *State) LoadTool(agentID, name string) error {
|
||||
if agentID == "" {
|
||||
return fmt.Errorf("agent ID required")
|
||||
}
|
||||
st.mu.RLock()
|
||||
reg := st.registry
|
||||
aid := st.agentID()
|
||||
st.mu.RUnlock()
|
||||
|
||||
if reg == nil {
|
||||
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
|
||||
}
|
||||
if st.OnToolsChanged != nil {
|
||||
|
|
@ -155,18 +154,20 @@ func (st *State) LoadTool(name string) error {
|
|||
return nil
|
||||
}
|
||||
|
||||
// UnloadTool removes a tool from the current agent's registry.
|
||||
func (st *State) UnloadTool(name string) error {
|
||||
// UnloadTool removes a tool from the given agent's registry.
|
||||
func (st *State) UnloadTool(agentID, name string) error {
|
||||
if agentID == "" {
|
||||
return fmt.Errorf("agent ID required")
|
||||
}
|
||||
st.mu.RLock()
|
||||
reg := st.registry
|
||||
aid := st.agentID()
|
||||
st.mu.RUnlock()
|
||||
|
||||
if reg == nil {
|
||||
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
|
||||
}
|
||||
if st.OnToolsChanged != nil {
|
||||
|
|
@ -208,7 +209,6 @@ func (st *State) NewProc(ctx context.Context, payload string, background bool) (
|
|||
// Look up tool
|
||||
st.mu.RLock()
|
||||
reg := st.registry
|
||||
aid := st.agentID()
|
||||
cwd := st.cwd
|
||||
yolo := st.yolo
|
||||
envCopy := make(map[string]string)
|
||||
|
|
@ -217,6 +217,12 @@ func (st *State) NewProc(ctx context.Context, payload string, background bool) (
|
|||
}
|
||||
st.mu.RUnlock()
|
||||
|
||||
aid := args["agent"]
|
||||
if aid == "" {
|
||||
return "", 0, fmt.Errorf("missing 'agent' in payload")
|
||||
}
|
||||
envCopy["OLLIE_UNAME"] = aid
|
||||
|
||||
if reg == nil {
|
||||
return "", 0, fmt.Errorf("tool registry not configured")
|
||||
}
|
||||
|
|
@ -456,15 +462,15 @@ func (st *State) HandleCtl(input string) error {
|
|||
|
||||
switch parts[0] {
|
||||
case "load":
|
||||
if len(parts) < 2 {
|
||||
return fmt.Errorf("load requires tool name")
|
||||
if len(parts) < 3 {
|
||||
return fmt.Errorf("load requires agent ID and tool name")
|
||||
}
|
||||
return st.LoadTool(parts[1])
|
||||
return st.LoadTool(parts[1], parts[2])
|
||||
case "unload":
|
||||
if len(parts) < 2 {
|
||||
return fmt.Errorf("unload requires tool name")
|
||||
if len(parts) < 3 {
|
||||
return fmt.Errorf("unload requires agent ID and tool name")
|
||||
}
|
||||
return st.UnloadTool(parts[1])
|
||||
return st.UnloadTool(parts[1], parts[2])
|
||||
case "env":
|
||||
if len(parts) < 2 {
|
||||
return fmt.Errorf("env requires KEY=VALUE")
|
||||
|
|
@ -475,7 +481,11 @@ func (st *State) HandleCtl(input string) error {
|
|||
if idx < 0 {
|
||||
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
|
||||
case "cwd":
|
||||
if len(parts) < 2 {
|
||||
|
|
@ -488,18 +498,13 @@ func (st *State) HandleCtl(input string) error {
|
|||
}
|
||||
}
|
||||
|
||||
// HandleToolsRead returns the tool list as JSON.
|
||||
func (st *State) HandleToolsRead() string {
|
||||
tools := st.ListTools()
|
||||
// HandleToolsRequest handles the tools rdwr: write agent ID, read tool list.
|
||||
func (st *State) HandleToolsRequest(agentID string) string {
|
||||
tools := st.ListTools(agentID)
|
||||
data, _ := json.Marshal(tools)
|
||||
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).
|
||||
func (st *State) HandleProcNew(ctx context.Context, payload string) (string, error) {
|
||||
result, _, err := st.NewProc(ctx, payload, false)
|
||||
|
|
|
|||
|
|
@ -37,13 +37,13 @@ func TestState_ToolsWithoutRegistry(t *testing.T) {
|
|||
st := NewState("/tmp")
|
||||
|
||||
// Without registry, ListTools should return nil
|
||||
tools := st.ListTools()
|
||||
tools := st.ListTools("agent1")
|
||||
if tools != nil {
|
||||
t.Errorf("ListTools() without registry = %v, want nil", tools)
|
||||
}
|
||||
|
||||
// Load should fail without registry
|
||||
err := st.LoadTool("shell")
|
||||
err := st.LoadTool("agent1", "shell")
|
||||
if err == nil {
|
||||
t.Error("LoadTool without registry should fail")
|
||||
}
|
||||
|
|
@ -58,16 +58,22 @@ func TestState_HandleCtl(t *testing.T) {
|
|||
t.Error("HandleCtl(unknown) should fail")
|
||||
}
|
||||
|
||||
// load without name
|
||||
// load without agent and name
|
||||
err = st.HandleCtl("load")
|
||||
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")
|
||||
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),
|
||||
),
|
||||
FileNode("tools", 0666,
|
||||
Doc("Tools: read to list, write tool name to load"),
|
||||
Read(handleToolsRead),
|
||||
Write(handleToolsWrite),
|
||||
Doc("Tools: write agent ID, read tool list (rdwr)"),
|
||||
Request(handleToolsRequest),
|
||||
),
|
||||
FileNode("info", 0444,
|
||||
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)))
|
||||
}
|
||||
|
||||
func handleToolsRead(ctx Ctx) ([]byte, error) {
|
||||
return []byte(ctx.Server.Fs.HandleToolsRead()), nil
|
||||
}
|
||||
|
||||
func handleToolsWrite(ctx Ctx, data []byte) error {
|
||||
return ctx.Server.Fs.HandleToolsWrite(strings.TrimSpace(string(data)))
|
||||
func handleToolsRequest(ctx Ctx, data []byte) ([]byte, error) {
|
||||
agentID := strings.TrimSpace(string(data))
|
||||
if agentID == "" {
|
||||
return nil, fmt.Errorf("agent ID required")
|
||||
}
|
||||
return []byte(ctx.Server.Fs.HandleToolsRequest(agentID)), nil
|
||||
}
|
||||
|
||||
func handleInfoRead(ctx Ctx) ([]byte, error) {
|
||||
|
|
|
|||
|
|
@ -20,9 +20,15 @@ type Conn struct {
|
|||
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 *Conn) 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).
|
||||
|
|
@ -95,21 +101,36 @@ func (c *Conn) Token() string {
|
|||
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) {
|
||||
fid, err := c.fsys.Open("tools", plan9.OREAD)
|
||||
fid, err := c.fsys.Open("tools", plan9.ORDWR)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer fid.Close()
|
||||
|
||||
data, err := io.ReadAll(fid)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
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 []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 tools, nil
|
||||
|
|
@ -129,7 +150,7 @@ func (c *Conn) CallTool(ctx context.Context, name string, args json.RawMessage)
|
|||
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 {
|
||||
escaped := escapeValue(fmt.Sprintf("%v", v))
|
||||
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{})
|
||||
}
|
||||
|
||||
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 {
|
||||
escaped := escapeValue(fmt.Sprintf("%v", v))
|
||||
payload += fmt.Sprintf("%s=%s\n", k, escaped)
|
||||
|
|
@ -203,7 +224,7 @@ func (c *Conn) LoadTool(name string) error {
|
|||
return err
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte("load " + name + "\n"))
|
||||
_, err = fid.Write([]byte("load " + c.agentID + " " + name + "\n"))
|
||||
return err
|
||||
}
|
||||
|
||||
|
|
@ -214,7 +235,7 @@ func (c *Conn) UnloadTool(name string) error {
|
|||
return err
|
||||
}
|
||||
defer fid.Close()
|
||||
_, err = fid.Write([]byte("unload " + name + "\n"))
|
||||
_, err = fid.Write([]byte("unload " + c.agentID + " " + name + "\n"))
|
||||
return err
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue