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:
parent
d6bc89d3e2
commit
10e0731a0a
|
|
@ -29,7 +29,7 @@ type Agent struct {
|
|||
agentsDir string
|
||||
baseLayers []string // system prompt layers for /agent reloads
|
||||
promptEnvExtra []string // PRIME_* vars for prompt resolution
|
||||
newToolServer func() tools.Server
|
||||
newToolServer func() tools.Runner
|
||||
newBackend func(string) (backend.Backend, error)
|
||||
currentAction atomic.Pointer[actionHandle]
|
||||
warnedContext bool
|
||||
|
|
@ -262,7 +262,7 @@ func (ag *Agent) SetSessionEnv(sessionID string) {
|
|||
return
|
||||
}
|
||||
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)
|
||||
if ag.id != "" {
|
||||
es.SetEnv("OLLIE_UNAME", ag.id)
|
||||
|
|
@ -277,7 +277,7 @@ func (ag *Agent) SetEnv(key, value string) {
|
|||
return
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
|
@ -304,7 +304,7 @@ func (ag *Agent) SetCWD(dir string) {
|
|||
}
|
||||
if ag.runtime.ExecServer != 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)
|
||||
}
|
||||
}
|
||||
|
|
@ -473,7 +473,7 @@ func (ag *Agent) React(responseID, emoji string) error {
|
|||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ import (
|
|||
// env provides additional environment variables injected into prompt resolution
|
||||
// subprocesses (e.g. OLLIE_SESSION_ID=xxx).
|
||||
// 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 allToolInfos []tools.ToolInfo
|
||||
|
|
@ -69,7 +69,7 @@ func BuildRuntime(cfg *AgentConfig, srv tools.Server, cwd string, env []string,
|
|||
}
|
||||
maxSteps = cfg.MaxSteps
|
||||
if len(cfg.AllowTools) > 0 {
|
||||
if rs, ok := srv.(tools.ToolRestrictionSetter); ok {
|
||||
if rs, ok := srv.(interface{ SetAllowTools([]string) }); ok {
|
||||
rs.SetAllowTools(cfg.AllowTools)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ type AgentCfg struct {
|
|||
CWD string // working directory for tool execution
|
||||
BaseLayers []string
|
||||
PromptEnvExtra []string
|
||||
NewToolServer func() tools.Server
|
||||
NewToolServer func() tools.Runner
|
||||
NewBackend func(string) (backend.Backend, error)
|
||||
Bus *pubsub.Bus
|
||||
Log *olog.Logger
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import (
|
|||
// agents replaces it atomically.
|
||||
type Runtime struct {
|
||||
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
|
||||
Preamble string // compiled system prompt
|
||||
Tools []backend.Tool
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ type Config struct {
|
|||
CWD string
|
||||
History *agent.History
|
||||
Runtime *agent.Runtime
|
||||
NewToolServer func() tools.Server
|
||||
NewToolServer func() tools.Runner
|
||||
NewBackend func(string) (backend.Backend, error)
|
||||
Log *olog.Logger
|
||||
MaxSteps int
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
package execute
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
// Package remote provides the SSH bootstrap and JSON-RPC client for
|
||||
// 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.
|
||||
package execute
|
||||
package tools
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
|
|
@ -25,7 +25,6 @@ import (
|
|||
"syscall"
|
||||
"time"
|
||||
|
||||
"ollie/tools"
|
||||
)
|
||||
|
||||
//go:embed bootstrap.sh
|
||||
|
|
@ -38,7 +37,7 @@ type HostInfo struct {
|
|||
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.
|
||||
type RemoteServer struct {
|
||||
mu sync.Mutex
|
||||
|
|
@ -253,9 +252,9 @@ func (s *RemoteServer) Close() error {
|
|||
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)
|
||||
req := rpcRequest{
|
||||
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)
|
||||
}
|
||||
|
||||
var infos []tools.ToolInfo
|
||||
var infos []ToolInfo
|
||||
if err := json.Unmarshal(resp.Result, &infos); err != nil {
|
||||
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.
|
||||
var notif outputNotification
|
||||
if json.Unmarshal(resp.Result, ¬if) == nil && notif.Data != "" {
|
||||
tools.StreamOutput(ctx, notif.Data)
|
||||
StreamOutput(ctx, notif.Data)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
|
@ -422,15 +421,15 @@ func shellEscape(s string) string {
|
|||
return "'" + strings.ReplaceAll(s, "'", "'\\''") + "'"
|
||||
}
|
||||
|
||||
// Decl returns a factory function compatible with tools.NewDispatcherFunc.
|
||||
// It dials the remote on first call and returns the Server.
|
||||
func RemoteDecl(cfg RemoteConfig) func() tools.Server {
|
||||
// Decl returns a factory function for remote tool servers.
|
||||
// It dials the remote on first call and returns a Runner.
|
||||
func RemoteDecl(cfg RemoteConfig) func() Runner {
|
||||
var (
|
||||
once sync.Once
|
||||
server *RemoteServer
|
||||
err error
|
||||
)
|
||||
return func() tools.Server {
|
||||
return func() Runner {
|
||||
once.Do(func() {
|
||||
server, err = RemoteDial(context.Background(), cfg)
|
||||
if err != nil {
|
||||
|
|
@ -439,21 +438,21 @@ func RemoteDecl(cfg RemoteConfig) func() tools.Server {
|
|||
}
|
||||
})
|
||||
if server == nil {
|
||||
return &errServer{err: err}
|
||||
return &errRunner{err: err}
|
||||
}
|
||||
return server
|
||||
}
|
||||
}
|
||||
|
||||
// errServer is a tools.Server that returns an error for every call.
|
||||
type errServer struct {
|
||||
// errRunner is a Runner that returns an error for every call.
|
||||
type errRunner struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (e *errServer) ListTools() ([]tools.ToolInfo, error) {
|
||||
func (e *errRunner) ListTools() ([]ToolInfo, error) {
|
||||
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
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package execute
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
|
@ -19,7 +19,6 @@ import (
|
|||
|
||||
"ollie/sandbox"
|
||||
"ollie/paths"
|
||||
"ollie/tools"
|
||||
"ollie/skills"
|
||||
"ollie/detach"
|
||||
)
|
||||
|
|
@ -56,7 +55,7 @@ type Server struct {
|
|||
// Empty means all are allowed.
|
||||
allowTools map[string]bool
|
||||
|
||||
toolRegistry *tools.Registry
|
||||
toolRegistry *Registry
|
||||
skillsRegistry *skills.Registry
|
||||
sessionID string
|
||||
|
||||
|
|
@ -125,7 +124,7 @@ func (e *Server) AllowTools() []string {
|
|||
|
||||
|
||||
// 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) {
|
||||
s.toolRegistry = r
|
||||
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.
|
||||
func Decl(cwd string, opts ...Option) func() tools.Server {
|
||||
return func() tools.Server {
|
||||
func Decl(cwd string, opts ...Option) func() Runner {
|
||||
return func() Runner {
|
||||
s := New(cwd)
|
||||
for _, o := range opts {
|
||||
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.
|
||||
func (e *Server) ListTools() ([]tools.ToolInfo, error) {
|
||||
all := []tools.ToolInfo{
|
||||
func (e *Server) ListTools() ([]ToolInfo, error) {
|
||||
all := []ToolInfo{
|
||||
{
|
||||
Name: "shell",
|
||||
Description: `Execute a single bash command in a sandboxed environment.
|
||||
|
|
@ -220,7 +219,7 @@ Returns tools with descriptions, one per line.`,
|
|||
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) {
|
||||
if e.toolRegistry != nil && e.sessionID != "" {
|
||||
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.
|
||||
func (e *Server) SetToolRegistry(r *tools.Registry, sessionID string) {
|
||||
func (e *Server) SetToolRegistry(r *Registry, sessionID string) {
|
||||
e.toolRegistry = r
|
||||
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, "..") {
|
||||
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).
|
||||
elevated := false
|
||||
|
|
@ -505,7 +504,7 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
|
|||
default:
|
||||
// Normal (foreground) execution — must also handle manual detach.
|
||||
var outputBuf bytes.Buffer
|
||||
streamFn := tools.StreamFunc(ctx)
|
||||
streamFn := StreamFunc(ctx)
|
||||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
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 = append(cmd.Env, "OLLIE_TOOLS_PATH="+tools.ToolsPath())
|
||||
cmd.Env = append(cmd.Env, "OLLIE_TOOLS_PATH="+ToolsPath())
|
||||
for k, v := range e.envExtra {
|
||||
cmd.Env = append(cmd.Env, k+"="+v)
|
||||
}
|
||||
|
|
@ -755,7 +754,7 @@ func (e *Server) executeWithStdin(ctx context.Context, code, language string, ti
|
|||
lw := &limitedWriter{
|
||||
w: &outputBuf,
|
||||
limit: 10 * 1024 * 1024,
|
||||
stream: tools.StreamFunc(ctx),
|
||||
stream: StreamFunc(ctx),
|
||||
}
|
||||
cmd.Stdout = lw
|
||||
cmd.Stderr = lw
|
||||
|
|
@ -1026,4 +1025,3 @@ func (e *Server) cleanupDetached() {
|
|||
}
|
||||
}
|
||||
|
||||
var _ tools.Server = (*Server)(nil) // compile-time interface check
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package execute
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
|
@ -7,7 +7,6 @@ import (
|
|||
"strings"
|
||||
|
||||
"ollie/skills"
|
||||
"ollie/tools"
|
||||
)
|
||||
|
||||
// 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.
|
||||
// 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 {
|
||||
return nil
|
||||
}
|
||||
tools := []tools.ToolInfo{
|
||||
tools := []ToolInfo{
|
||||
{
|
||||
Name: "skill_list",
|
||||
Description: `List available skill modules with name and description.
|
||||
|
|
@ -1,8 +1,7 @@
|
|||
package execute
|
||||
package tools
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"ollie/tools"
|
||||
"strings"
|
||||
)
|
||||
|
||||
|
|
@ -14,7 +13,7 @@ func (e *Server) ResultTier(name string) string {
|
|||
return info.Tier
|
||||
}
|
||||
}
|
||||
code, err := tools.ReadTool(name)
|
||||
code, err := ReadTool(name)
|
||||
if err != nil {
|
||||
return "hot"
|
||||
}
|
||||
|
|
@ -42,7 +41,7 @@ func (e *Server) IsParallelRead(name string) bool {
|
|||
return info.ReadOnly
|
||||
}
|
||||
}
|
||||
code, err := tools.ReadTool(name)
|
||||
code, err := ReadTool(name)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
|
@ -1,11 +1,10 @@
|
|||
// Package tools defines the Server and Dispatcher interfaces and their
|
||||
// default implementations.
|
||||
// Package tools implements the tool server (sandboxed execution, tool registry,
|
||||
// skill management) and supporting types.
|
||||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// ToolInfo describes a tool provided by a server.
|
||||
|
|
@ -24,97 +23,13 @@ type ToolInfo struct {
|
|||
ReadOnly bool
|
||||
}
|
||||
|
||||
// Server is the interface satisfied by any tool server.
|
||||
type Server interface {
|
||||
// Runner is the minimal interface satisfied by any tool server (local or remote).
|
||||
// Consumers that need polymorphism over Server and RemoteServer use this.
|
||||
type Runner interface {
|
||||
ListTools() ([]ToolInfo, 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
|
||||
// named tool is safe to run concurrently with other read-class tools.
|
||||
// Returns false for unknown tools (conservative default).
|
||||
|
|
|
|||
|
|
@ -3,35 +3,34 @@ package tools_test
|
|||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"ollie/tools"
|
||||
)
|
||||
|
||||
// stubServer is a minimal Server used to verify the contract.
|
||||
type stubServer struct {
|
||||
// stubRunner is a minimal Runner used to verify the contract.
|
||||
type stubRunner struct {
|
||||
name string
|
||||
tools []tools.ToolInfo
|
||||
}
|
||||
|
||||
func (s *stubServer) ListTools() ([]tools.ToolInfo, error) { return s.tools, nil }
|
||||
func (s *stubServer) CallTool(_ context.Context, tool string, _ json.RawMessage) (json.RawMessage, error) {
|
||||
func (s *stubRunner) ListTools() ([]tools.ToolInfo, error) { return s.tools, nil }
|
||||
func (s *stubRunner) CallTool(_ context.Context, tool string, _ json.RawMessage) (json.RawMessage, error) {
|
||||
return json.RawMessage(`{"tool":"` + tool + `"}`), nil
|
||||
}
|
||||
|
||||
func newStub(name string, toolNames ...string) *stubServer {
|
||||
func newStub(name string, toolNames ...string) *stubRunner {
|
||||
var ti []tools.ToolInfo
|
||||
for _, n := range toolNames {
|
||||
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.
|
||||
func checkServerContract(t *testing.T, s tools.Server) {
|
||||
// checkRunnerContract verifies Runner invariants.
|
||||
func checkRunnerContract(t *testing.T, r tools.Runner) {
|
||||
t.Helper()
|
||||
tl, err := s.ListTools()
|
||||
tl, err := r.ListTools()
|
||||
if err != nil {
|
||||
t.Fatalf("ListTools: %v", err)
|
||||
}
|
||||
|
|
@ -39,7 +38,7 @@ func checkServerContract(t *testing.T, s tools.Server) {
|
|||
t.Fatal("ListTools returned nil")
|
||||
}
|
||||
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 {
|
||||
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 checkDispatcherContract(t *testing.T, d tools.Dispatcher, servers map[string]*stubServer) {
|
||||
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))
|
||||
}
|
||||
func TestStubRunnerContract(t *testing.T) {
|
||||
checkRunnerContract(t, newStub("s", "a", "b"))
|
||||
}
|
||||
|
|
|
|||
Reference in New Issue