merge execute/ into tools/

tools.Server is now a concrete struct (the local execution engine).
tools.Runner is the minimal 2-method interface for polymorphism
(satisfied by both Server and RemoteServer).

Deleted: Dispatcher, CWDSetter, EnvSetter, ToolRestrictionSetter
interfaces. Agent uses inline type assertions where needed.

The execute/ package no longer exists.
This commit is contained in:
Levi Neely 2026-07-29 23:01:47 +02:00
parent d6bc89d3e2
commit 10e0731a0a
13 changed files with 64 additions and 250 deletions

View File

@ -29,7 +29,7 @@ type Agent struct {
agentsDir string agentsDir string
baseLayers []string // system prompt layers for /agent reloads baseLayers []string // system prompt layers for /agent reloads
promptEnvExtra []string // PRIME_* vars for prompt resolution promptEnvExtra []string // PRIME_* vars for prompt resolution
newToolServer func() tools.Server newToolServer func() tools.Runner
newBackend func(string) (backend.Backend, error) newBackend func(string) (backend.Backend, error)
currentAction atomic.Pointer[actionHandle] currentAction atomic.Pointer[actionHandle]
warnedContext bool warnedContext bool
@ -262,7 +262,7 @@ func (ag *Agent) SetSessionEnv(sessionID string) {
return return
} }
if srv := ag.runtime.ExecServer; srv != nil { if srv := ag.runtime.ExecServer; srv != nil {
if es, ok := srv.(tools.EnvSetter); ok { if es, ok := srv.(interface{ SetEnv(string, string) }); ok {
es.SetEnv("OLLIE_SESSION_ID", sessionID) es.SetEnv("OLLIE_SESSION_ID", sessionID)
if ag.id != "" { if ag.id != "" {
es.SetEnv("OLLIE_UNAME", ag.id) es.SetEnv("OLLIE_UNAME", ag.id)
@ -277,7 +277,7 @@ func (ag *Agent) SetEnv(key, value string) {
return return
} }
if srv := ag.runtime.ExecServer; srv != nil { if srv := ag.runtime.ExecServer; srv != nil {
if es, ok := srv.(tools.EnvSetter); ok { if es, ok := srv.(interface{ SetEnv(string, string) }); ok {
es.SetEnv(key, value) es.SetEnv(key, value)
} }
} }
@ -304,7 +304,7 @@ func (ag *Agent) SetCWD(dir string) {
} }
if ag.runtime.ExecServer != nil { if ag.runtime.ExecServer != nil {
if srv := ag.runtime.ExecServer; srv != nil { if srv := ag.runtime.ExecServer; srv != nil {
if ws, ok := srv.(tools.CWDSetter); ok { if ws, ok := srv.(interface{ SetCWD(string) }); ok {
ws.SetCWD(dir) ws.SetCWD(dir)
} }
} }
@ -473,7 +473,7 @@ func (ag *Agent) React(responseID, emoji string) error {
} }
// execServer returns the execute server, or nil if unavailable. // execServer returns the execute server, or nil if unavailable.
func (ag *Agent) execServer() tools.Server { func (ag *Agent) execServer() tools.Runner {
return ag.runtime.ExecServer return ag.runtime.ExecServer
} }

View File

@ -17,7 +17,7 @@ import (
// env provides additional environment variables injected into prompt resolution // env provides additional environment variables injected into prompt resolution
// subprocesses (e.g. OLLIE_SESSION_ID=xxx). // subprocesses (e.g. OLLIE_SESSION_ID=xxx).
// The caller is responsible for registering all servers on d before calling this. // The caller is responsible for registering all servers on d before calling this.
func BuildRuntime(cfg *AgentConfig, srv tools.Server, cwd string, env []string, baseLayers ...string) *Runtime { func BuildRuntime(cfg *AgentConfig, srv tools.Runner, cwd string, env []string, baseLayers ...string) *Runtime {
var messages []string var messages []string
var allToolInfos []tools.ToolInfo var allToolInfos []tools.ToolInfo
@ -69,7 +69,7 @@ func BuildRuntime(cfg *AgentConfig, srv tools.Server, cwd string, env []string,
} }
maxSteps = cfg.MaxSteps maxSteps = cfg.MaxSteps
if len(cfg.AllowTools) > 0 { if len(cfg.AllowTools) > 0 {
if rs, ok := srv.(tools.ToolRestrictionSetter); ok { if rs, ok := srv.(interface{ SetAllowTools([]string) }); ok {
rs.SetAllowTools(cfg.AllowTools) rs.SetAllowTools(cfg.AllowTools)
} }
} }

View File

@ -19,7 +19,7 @@ type AgentCfg struct {
CWD string // working directory for tool execution CWD string // working directory for tool execution
BaseLayers []string BaseLayers []string
PromptEnvExtra []string PromptEnvExtra []string
NewToolServer func() tools.Server NewToolServer func() tools.Runner
NewBackend func(string) (backend.Backend, error) NewBackend func(string) (backend.Backend, error)
Bus *pubsub.Bus Bus *pubsub.Bus
Log *olog.Logger Log *olog.Logger

View File

@ -13,7 +13,7 @@ import (
// agents replaces it atomically. // agents replaces it atomically.
type Runtime struct { type Runtime struct {
Backend backend.Backend Backend backend.Backend
ExecServer tools.Server // the execute server (tool runtime, env, cwd) ExecServer tools.Runner // the execute server (tool runtime, env, cwd)
Hooks Hooks Hooks Hooks
Preamble string // compiled system prompt Preamble string // compiled system prompt
Tools []backend.Tool Tools []backend.Tool

View File

@ -32,7 +32,7 @@ type Config struct {
CWD string CWD string
History *agent.History History *agent.History
Runtime *agent.Runtime Runtime *agent.Runtime
NewToolServer func() tools.Server NewToolServer func() tools.Runner
NewBackend func(string) (backend.Backend, error) NewBackend func(string) (backend.Backend, error)
Log *olog.Logger Log *olog.Logger
MaxSteps int MaxSteps int

View File

@ -1,4 +1,4 @@
package execute package tools
import ( import (
"context" "context"

View File

@ -1,8 +1,8 @@
// Package remote provides the SSH bootstrap and JSON-RPC client for // Package remote provides the SSH bootstrap and JSON-RPC client for
// split-brain remote execution. It connects to a remote host over SSH, // split-brain remote execution. It connects to a remote host over SSH,
// ensures ollie-remote is deployed, and returns a tools.Server that // ensures ollie-remote is deployed, and returns a Server that
// forwards execution calls over the RPC channel. // forwards execution calls over the RPC channel.
package execute package tools
import ( import (
"bufio" "bufio"
@ -25,7 +25,6 @@ import (
"syscall" "syscall"
"time" "time"
"ollie/tools"
) )
//go:embed bootstrap.sh //go:embed bootstrap.sh
@ -38,7 +37,7 @@ type HostInfo struct {
IsGitRepo bool `json:"is_git_repo"` IsGitRepo bool `json:"is_git_repo"`
} }
// Server implements tools.Server by forwarding calls to a remote // Server implements Server by forwarding calls to a remote
// ollie-remote process over SSH. // ollie-remote process over SSH.
type RemoteServer struct { type RemoteServer struct {
mu sync.Mutex mu sync.Mutex
@ -253,9 +252,9 @@ func (s *RemoteServer) Close() error {
return nil return nil
} }
// --- tools.Server interface --- // --- Server interface ---
func (s *RemoteServer) ListTools() ([]tools.ToolInfo, error) { func (s *RemoteServer) ListTools() ([]ToolInfo, error) {
id := s.nextID.Add(1) id := s.nextID.Add(1)
req := rpcRequest{ req := rpcRequest{
JSONRPC: "2.0", JSONRPC: "2.0",
@ -279,7 +278,7 @@ func (s *RemoteServer) ListTools() ([]tools.ToolInfo, error) {
return nil, fmt.Errorf("remote list_tools: %s", resp.Error.Message) return nil, fmt.Errorf("remote list_tools: %s", resp.Error.Message)
} }
var infos []tools.ToolInfo var infos []ToolInfo
if err := json.Unmarshal(resp.Result, &infos); err != nil { if err := json.Unmarshal(resp.Result, &infos); err != nil {
return nil, fmt.Errorf("remote list_tools unmarshal: %w", err) return nil, fmt.Errorf("remote list_tools unmarshal: %w", err)
} }
@ -317,7 +316,7 @@ func (s *RemoteServer) CallTool(ctx context.Context, tool string, args json.RawM
// Streaming output notification — emit to context callback. // Streaming output notification — emit to context callback.
var notif outputNotification var notif outputNotification
if json.Unmarshal(resp.Result, &notif) == nil && notif.Data != "" { if json.Unmarshal(resp.Result, &notif) == nil && notif.Data != "" {
tools.StreamOutput(ctx, notif.Data) StreamOutput(ctx, notif.Data)
} }
continue continue
} }
@ -422,15 +421,15 @@ func shellEscape(s string) string {
return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'" return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'"
} }
// Decl returns a factory function compatible with tools.NewDispatcherFunc. // Decl returns a factory function for remote tool servers.
// It dials the remote on first call and returns the Server. // It dials the remote on first call and returns a Runner.
func RemoteDecl(cfg RemoteConfig) func() tools.Server { func RemoteDecl(cfg RemoteConfig) func() Runner {
var ( var (
once sync.Once once sync.Once
server *RemoteServer server *RemoteServer
err error err error
) )
return func() tools.Server { return func() Runner {
once.Do(func() { once.Do(func() {
server, err = RemoteDial(context.Background(), cfg) server, err = RemoteDial(context.Background(), cfg)
if err != nil { if err != nil {
@ -439,21 +438,21 @@ func RemoteDecl(cfg RemoteConfig) func() tools.Server {
} }
}) })
if server == nil { if server == nil {
return &errServer{err: err} return &errRunner{err: err}
} }
return server return server
} }
} }
// errServer is a tools.Server that returns an error for every call. // errRunner is a Runner that returns an error for every call.
type errServer struct { type errRunner struct {
err error err error
} }
func (e *errServer) ListTools() ([]tools.ToolInfo, error) { func (e *errRunner) ListTools() ([]ToolInfo, error) {
return nil, e.err return nil, e.err
} }
func (e *errServer) CallTool(_ context.Context, _ string, _ json.RawMessage) (json.RawMessage, error) { func (e *errRunner) CallTool(_ context.Context, _ string, _ json.RawMessage) (json.RawMessage, error) {
return nil, e.err return nil, e.err
} }

View File

@ -1,4 +1,4 @@
package execute package tools
import ( import (
"context" "context"
@ -19,7 +19,6 @@ import (
"ollie/sandbox" "ollie/sandbox"
"ollie/paths" "ollie/paths"
"ollie/tools"
"ollie/skills" "ollie/skills"
"ollie/detach" "ollie/detach"
) )
@ -56,7 +55,7 @@ type Server struct {
// Empty means all are allowed. // Empty means all are allowed.
allowTools map[string]bool allowTools map[string]bool
toolRegistry *tools.Registry toolRegistry *Registry
skillsRegistry *skills.Registry skillsRegistry *skills.Registry
sessionID string sessionID string
@ -125,7 +124,7 @@ func (e *Server) AllowTools() []string {
// WithToolRegistry attaches a tool registry and session ID to the Server. // WithToolRegistry attaches a tool registry and session ID to the Server.
func WithToolRegistry(r *tools.Registry, sessionID string) Option { func WithToolRegistry(r *Registry, sessionID string) Option {
return func(s *Server) { return func(s *Server) {
s.toolRegistry = r s.toolRegistry = r
s.sessionID = sessionID s.sessionID = sessionID
@ -138,8 +137,8 @@ func WithSkillsRegistry(r *skills.Registry) Option {
} }
// Decl returns a factory for an execute Server with the given working directory. // Decl returns a factory for an execute Server with the given working directory.
func Decl(cwd string, opts ...Option) func() tools.Server { func Decl(cwd string, opts ...Option) func() Runner {
return func() tools.Server { return func() Runner {
s := New(cwd) s := New(cwd)
for _, o := range opts { for _, o := range opts {
o(s) o(s)
@ -148,10 +147,10 @@ func Decl(cwd string, opts ...Option) func() tools.Server {
} }
} }
// ListTools implements tools.Server, returning shell plus any // ListTools implements Server, returning shell plus any
// tools promoted in the session's tool registry. // tools promoted in the session's tool registry.
func (e *Server) ListTools() ([]tools.ToolInfo, error) { func (e *Server) ListTools() ([]ToolInfo, error) {
all := []tools.ToolInfo{ all := []ToolInfo{
{ {
Name: "shell", Name: "shell",
Description: `Execute a single bash command in a sandboxed environment. Description: `Execute a single bash command in a sandboxed environment.
@ -220,7 +219,7 @@ Returns tools with descriptions, one per line.`,
return all, nil return all, nil
} }
// CallTool implements tools.Server. // CallTool implements Server.
func (e *Server) CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) { func (e *Server) CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) {
if e.toolRegistry != nil && e.sessionID != "" { if e.toolRegistry != nil && e.sessionID != "" {
if _, promoted := e.toolRegistry.Lookup(e.sessionID, tool); promoted { if _, promoted := e.toolRegistry.Lookup(e.sessionID, tool); promoted {
@ -283,7 +282,7 @@ func (e *Server) SetEnv(key, value string) {
} }
// SetToolRegistry attaches a session-local tool registry. // SetToolRegistry attaches a session-local tool registry.
func (e *Server) SetToolRegistry(r *tools.Registry, sessionID string) { func (e *Server) SetToolRegistry(r *Registry, sessionID string) {
e.toolRegistry = r e.toolRegistry = r
e.sessionID = sessionID e.sessionID = sessionID
} }
@ -295,7 +294,7 @@ func (e *Server) callPromotedTool(ctx context.Context, tool string, args json.Ra
if strings.Contains(tool, "/") || strings.Contains(tool, "..") { if strings.Contains(tool, "/") || strings.Contains(tool, "..") {
return nil, fmt.Errorf("invalid tool name") return nil, fmt.Errorf("invalid tool name")
} }
path := filepath.Join(tools.ToolsPath(), tool) path := filepath.Join(ToolsPath(), tool)
// Extract elevated flag (dispatch-level concern, not passed to tool). // Extract elevated flag (dispatch-level concern, not passed to tool).
elevated := false elevated := false
@ -505,7 +504,7 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
default: default:
// Normal (foreground) execution — must also handle manual detach. // Normal (foreground) execution — must also handle manual detach.
var outputBuf bytes.Buffer var outputBuf bytes.Buffer
streamFn := tools.StreamFunc(ctx) streamFn := StreamFunc(ctx)
lw := &limitedWriter{ lw := &limitedWriter{
w: &outputBuf, w: &outputBuf,
limit: 10 * 1024 * 1024, limit: 10 * 1024 * 1024,
@ -736,7 +735,7 @@ func (e *Server) executeWithStdin(ctx context.Context, code, language string, ti
} }
} }
cmd.Env = prependOlliePath(filtered, paths.CfgDir()) cmd.Env = prependOlliePath(filtered, paths.CfgDir())
cmd.Env = append(cmd.Env, "OLLIE_TOOLS_PATH="+tools.ToolsPath()) cmd.Env = append(cmd.Env, "OLLIE_TOOLS_PATH="+ToolsPath())
for k, v := range e.envExtra { for k, v := range e.envExtra {
cmd.Env = append(cmd.Env, k+"="+v) cmd.Env = append(cmd.Env, k+"="+v)
} }
@ -755,7 +754,7 @@ func (e *Server) executeWithStdin(ctx context.Context, code, language string, ti
lw := &limitedWriter{ lw := &limitedWriter{
w: &outputBuf, w: &outputBuf,
limit: 10 * 1024 * 1024, limit: 10 * 1024 * 1024,
stream: tools.StreamFunc(ctx), stream: StreamFunc(ctx),
} }
cmd.Stdout = lw cmd.Stdout = lw
cmd.Stderr = lw cmd.Stderr = lw
@ -1026,4 +1025,3 @@ func (e *Server) cleanupDetached() {
} }
} }
var _ tools.Server = (*Server)(nil) // compile-time interface check

View File

@ -1,4 +1,4 @@
package execute package tools
import ( import (
"context" "context"
@ -7,7 +7,6 @@ import (
"strings" "strings"
"ollie/skills" "ollie/skills"
"ollie/tools"
) )
// SetSkillsRegistry attaches a skills registry to the execute server. // SetSkillsRegistry attaches a skills registry to the execute server.
@ -18,11 +17,11 @@ func (e *Server) SetSkillsRegistry(r *skills.Registry) {
// ListSkillsTools returns ToolInfo entries for the skill_* built-ins. // ListSkillsTools returns ToolInfo entries for the skill_* built-ins.
// These are included alongside the standard tool_* tools. // These are included alongside the standard tool_* tools.
func ListSkillsTools(skillsReg *skills.Registry, sessionID string) []tools.ToolInfo { func ListSkillsTools(skillsReg *skills.Registry, sessionID string) []ToolInfo {
if skillsReg == nil { if skillsReg == nil {
return nil return nil
} }
tools := []tools.ToolInfo{ tools := []ToolInfo{
{ {
Name: "skill_list", Name: "skill_list",
Description: `List available skill modules with name and description. Description: `List available skill modules with name and description.

View File

@ -1,8 +1,7 @@
package execute package tools
import ( import (
"encoding/json" "encoding/json"
"ollie/tools"
"strings" "strings"
) )
@ -14,7 +13,7 @@ func (e *Server) ResultTier(name string) string {
return info.Tier return info.Tier
} }
} }
code, err := tools.ReadTool(name) code, err := ReadTool(name)
if err != nil { if err != nil {
return "hot" return "hot"
} }
@ -42,7 +41,7 @@ func (e *Server) IsParallelRead(name string) bool {
return info.ReadOnly return info.ReadOnly
} }
} }
code, err := tools.ReadTool(name) code, err := ReadTool(name)
if err != nil { if err != nil {
return false return false
} }

View File

@ -1,11 +1,10 @@
// Package tools defines the Server and Dispatcher interfaces and their // Package tools implements the tool server (sandboxed execution, tool registry,
// default implementations. // skill management) and supporting types.
package tools package tools
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"fmt"
) )
// ToolInfo describes a tool provided by a server. // ToolInfo describes a tool provided by a server.
@ -24,97 +23,13 @@ type ToolInfo struct {
ReadOnly bool ReadOnly bool
} }
// Server is the interface satisfied by any tool server. // Runner is the minimal interface satisfied by any tool server (local or remote).
type Server interface { // Consumers that need polymorphism over Server and RemoteServer use this.
type Runner interface {
ListTools() ([]ToolInfo, error) ListTools() ([]ToolInfo, error)
CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error)
} }
// Dispatcher routes tool calls to the server that owns them.
type Dispatcher interface {
AddServer(name string, s Server)
GetServer(name string) (Server, bool)
ListTools() ([]ToolInfo, error)
Dispatch(ctx context.Context, server, tool string, args json.RawMessage) (json.RawMessage, error)
}
// dispatcher is the default Dispatcher backed by registered Server instances.
type dispatcher struct {
servers map[string]Server
}
// NewDispatcher returns a Dispatcher with no servers registered.
func NewDispatcher() Dispatcher {
return &dispatcher{servers: make(map[string]Server)}
}
// NewDispatcherFunc returns a factory that builds a fresh Dispatcher on each
// call, invoking each decl to create its Server. Pass the result to
// agent.AgentCoreConfig.NewDispatcher.
func NewDispatcherFunc(decls map[string]func() Server) func() Dispatcher {
return func() Dispatcher {
d := NewDispatcher()
for name, decl := range decls {
d.AddServer(name, decl())
}
return d
}
}
// AddServer registers a Server under the given name.
func (d *dispatcher) AddServer(name string, s Server) {
d.servers[name] = s
}
// GetServer returns the Server registered under the given name, if any.
func (d *dispatcher) GetServer(name string) (Server, bool) {
s, ok := d.servers[name]
return s, ok
}
// ListTools returns all tools advertised by all registered servers.
func (d *dispatcher) ListTools() ([]ToolInfo, error) {
var all []ToolInfo
for serverName, s := range d.servers {
tools, err := s.ListTools()
if err != nil {
return nil, fmt.Errorf("server %s: %w", serverName, err)
}
for _, t := range tools {
t.Server = serverName
all = append(all, t)
}
}
return all, nil
}
// Dispatch calls a named tool on the named server.
func (d *dispatcher) Dispatch(ctx context.Context, server, tool string, args json.RawMessage) (json.RawMessage, error) {
s, ok := d.servers[server]
if !ok {
return nil, fmt.Errorf("server not found: %s", server)
}
return s.CallTool(ctx, tool, args)
}
// CWDSetter is implemented by tool servers that accept a dynamic working
// directory. SetCWD updates the directory used for subsequent tool calls.
type CWDSetter interface {
SetCWD(string)
}
// EnvSetter is implemented by tool servers that accept per-session environment
// variables. SetEnv adds a key=value pair to the command environment.
type EnvSetter interface {
SetEnv(key, value string)
}
// ToolRestrictionSetter is implemented by tool servers that support restricting
// which tool scripts are available.
type ToolRestrictionSetter interface {
SetAllowTools(names []string)
}
// ParallelClassifier is implemented by tool servers that can report whether a // ParallelClassifier is implemented by tool servers that can report whether a
// named tool is safe to run concurrently with other read-class tools. // named tool is safe to run concurrently with other read-class tools.
// Returns false for unknown tools (conservative default). // Returns false for unknown tools (conservative default).

View File

@ -3,35 +3,34 @@ package tools_test
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"fmt"
"testing" "testing"
"ollie/tools" "ollie/tools"
) )
// stubServer is a minimal Server used to verify the contract. // stubRunner is a minimal Runner used to verify the contract.
type stubServer struct { type stubRunner struct {
name string name string
tools []tools.ToolInfo tools []tools.ToolInfo
} }
func (s *stubServer) ListTools() ([]tools.ToolInfo, error) { return s.tools, nil } func (s *stubRunner) ListTools() ([]tools.ToolInfo, error) { return s.tools, nil }
func (s *stubServer) CallTool(_ context.Context, tool string, _ json.RawMessage) (json.RawMessage, error) { func (s *stubRunner) CallTool(_ context.Context, tool string, _ json.RawMessage) (json.RawMessage, error) {
return json.RawMessage(`{"tool":"` + tool + `"}`), nil return json.RawMessage(`{"tool":"` + tool + `"}`), nil
} }
func newStub(name string, toolNames ...string) *stubServer { func newStub(name string, toolNames ...string) *stubRunner {
var ti []tools.ToolInfo var ti []tools.ToolInfo
for _, n := range toolNames { for _, n := range toolNames {
ti = append(ti, tools.ToolInfo{Name: n, Description: n + " desc"}) ti = append(ti, tools.ToolInfo{Name: n, Description: n + " desc"})
} }
return &stubServer{name: name, tools: ti} return &stubRunner{name: name, tools: ti}
} }
// checkServerContract verifies Server invariants. // checkRunnerContract verifies Runner invariants.
func checkServerContract(t *testing.T, s tools.Server) { func checkRunnerContract(t *testing.T, r tools.Runner) {
t.Helper() t.Helper()
tl, err := s.ListTools() tl, err := r.ListTools()
if err != nil { if err != nil {
t.Fatalf("ListTools: %v", err) t.Fatalf("ListTools: %v", err)
} }
@ -39,7 +38,7 @@ func checkServerContract(t *testing.T, s tools.Server) {
t.Fatal("ListTools returned nil") t.Fatal("ListTools returned nil")
} }
for _, ti := range tl { for _, ti := range tl {
res, err := s.CallTool(context.Background(), ti.Name, json.RawMessage(`{}`)) res, err := r.CallTool(context.Background(), ti.Name, json.RawMessage(`{}`))
if err != nil { if err != nil {
t.Errorf("CallTool(%q): %v", ti.Name, err) t.Errorf("CallTool(%q): %v", ti.Name, err)
} }
@ -49,101 +48,6 @@ func checkServerContract(t *testing.T, s tools.Server) {
} }
} }
// checkDispatcherContract verifies Dispatcher invariants. func TestStubRunnerContract(t *testing.T) {
func checkDispatcherContract(t *testing.T, d tools.Dispatcher, servers map[string]*stubServer) { checkRunnerContract(t, newStub("s", "a", "b"))
t.Helper()
for name, s := range servers {
d.AddServer(name, s)
}
// GetServer round-trip
for name := range servers {
s, ok := d.GetServer(name)
if !ok || s == nil {
t.Errorf("GetServer(%q) not found after AddServer", name)
}
}
if _, ok := d.GetServer("nonexistent"); ok {
t.Error("GetServer returned true for unregistered server")
}
// ListTools aggregates all servers
all, err := d.ListTools()
if err != nil {
t.Fatalf("ListTools: %v", err)
}
var wantCount int
for _, s := range servers {
wantCount += len(s.tools)
}
if len(all) != wantCount {
t.Errorf("ListTools returned %d tools, want %d", len(all), wantCount)
}
for _, ti := range all {
if ti.Server == "" {
t.Errorf("tool %q has empty Server field", ti.Name)
}
}
// Dispatch routes to correct server
for name, s := range servers {
for _, ti := range s.tools {
res, err := d.Dispatch(context.Background(), name, ti.Name, json.RawMessage(`{}`))
if err != nil {
t.Errorf("Dispatch(%q, %q): %v", name, ti.Name, err)
}
if len(res) == 0 {
t.Errorf("Dispatch(%q, %q) returned empty", name, ti.Name)
}
}
}
// Dispatch to unknown server must error
_, err = d.Dispatch(context.Background(), "nonexistent", "tool", json.RawMessage(`{}`))
if err == nil {
t.Error("Dispatch to unknown server should error")
}
}
func TestStubServerContract(t *testing.T) {
checkServerContract(t, newStub("s", "a", "b"))
}
func TestDispatcherContract(t *testing.T) {
servers := map[string]*stubServer{
"alpha": newStub("alpha", "tool1", "tool2"),
"beta": newStub("beta", "tool3"),
}
checkDispatcherContract(t, tools.NewDispatcher(), servers)
}
// failServer is a Server whose ListTools always errors.
type failServer struct{ stubServer }
func (f *failServer) ListTools() ([]tools.ToolInfo, error) {
return nil, fmt.Errorf("boom")
}
func TestDispatcherListToolsError(t *testing.T) {
d := tools.NewDispatcher()
d.AddServer("bad", &failServer{})
_, err := d.ListTools()
if err == nil {
t.Error("expected error from failing server")
}
}
func TestNewDispatcherFunc(t *testing.T) {
factory := tools.NewDispatcherFunc(map[string]func() tools.Server{
"s1": func() tools.Server { return newStub("s1", "t1") },
})
d := factory()
tl, err := d.ListTools()
if err != nil {
t.Fatalf("ListTools: %v", err)
}
if len(tl) != 1 {
t.Errorf("got %d tools, want 1", len(tl))
}
} }