tool registry: cache tier/parallel metadata in ToolInfo, drop hard-coded tables

This commit is contained in:
ollie 2026-07-28 23:15:52 +02:00
parent 2c36c023aa
commit ba9f605119
11 changed files with 58 additions and 1184 deletions

View File

@ -400,11 +400,7 @@ func (s *agent) pushLockDir() {
if s.runtime == nil || s.runtime.Dispatcher == nil || s.sessionID == "" {
return
}
if srv, ok := s.runtime.Dispatcher.GetServer("execute"); ok {
if ls, ok := srv.(tools.LockDirSetter); ok {
ls.SetLockDir(filepath.Join(ollieTmpDir(), s.sessionID))
}
}
}
var _ Core = (*agent)(nil) // compile-time interface check

View File

@ -1,131 +0,0 @@
package agent
import (
"context"
"encoding/json"
"testing"
"ollie/pkg/backend"
)
// TestTextToolCall_ParsedAndExecuted verifies that when a model emits a tool
// call as plain text (no API-level tool calls) the loop parses it and executes
// it as a normal tool call.
func TestTextToolCall_ParsedAndExecuted(t *testing.T) {
be := defaultBE()
be.respond = func(_ context.Context, msgs []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
// First call: model emits text tool call
for _, m := range msgs {
if m.Role == "tool" {
// Second call after tool result: respond normally
return textStream("done"), nil
}
}
return textStream(`call_tool: {"calls":[{"tool":"file_read","args":["/tmp/test.go"]}]}`), nil
}
var execCalled bool
c := newCore(t, be, nil)
c.runtime.Tools = []backend.Tool{
{Name: "call_tool", Description: "run tools"},
{Name: "execute_code", Description: "run code"},
}
c.runtime.Exec = func(_ context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error) {
execCalled = true
if name != "call_tool" {
t.Errorf("expected tool name 'call_tool', got %q", name)
}
return "file contents here", nil, nil
}
evs := collectEvents(context.Background(), c, "read a file")
if !execCalled {
t.Fatal("expected Exec to be called for text-based tool call")
}
errEvs := byRole(evs, "error")
if len(errEvs) != 0 {
t.Errorf("unexpected error events: %v", errEvs)
}
}
// TestTextToolCall_RelaxedJSON verifies parsing of relaxed JSON (unquoted keys).
func TestTextToolCall_RelaxedJSON(t *testing.T) {
be := defaultBE()
be.respond = func(_ context.Context, msgs []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
for _, m := range msgs {
if m.Role == "tool" {
return textStream("done"), nil
}
}
return textStream(`execute_code: steps=[{code: "date"}]`), nil
}
var execCalled bool
c := newCore(t, be, nil)
c.runtime.Tools = []backend.Tool{
{Name: "execute_code", Description: "run code"},
{Name: "call_tool", Description: "run tools"},
}
c.runtime.Exec = func(_ context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error) {
execCalled = true
if name != "execute_code" {
t.Errorf("expected tool name 'execute_code', got %q", name)
}
return "Sat May 23 11:00:00 UTC 2026", nil, nil
}
collectEvents(context.Background(), c, "what time is it")
if !execCalled {
t.Fatal("expected Exec to be called for relaxed-JSON text tool call")
}
}
// TestTextToolCall_NoFalsePositive verifies that a normal text response that
// does not contain a tool-call pattern is not mistakenly parsed.
func TestTextToolCall_NoFalsePositive(t *testing.T) {
be := defaultBE()
be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
return textStream("Here is how execute_code works in general."), nil
}
c := newCore(t, be, nil)
c.runtime.Tools = []backend.Tool{{Name: "execute_code", Description: "run code"}}
evs := collectEvents(context.Background(), c, "explain tools")
errEvs := byRole(evs, "error")
if len(errEvs) != 0 {
t.Errorf("unexpected error events: %v", errEvs)
}
assistEvs := byRole(evs, "assistant")
if len(assistEvs) == 0 {
t.Error("expected assistant response; got none")
}
}
// TestTextToolCall_NoToolsConfigured verifies that when no tools are configured
// the detection does not fire (nothing to detect against).
func TestTextToolCall_NoToolsConfigured(t *testing.T) {
be := defaultBE()
be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
return textStream(`execute_code: steps=[{code: "print('hi')"}]`), nil
}
c := newCore(t, be, nil)
// No tools configured.
evs := collectEvents(context.Background(), c, "run something")
errEvs := byRole(evs, "error")
if len(errEvs) != 0 {
t.Errorf("unexpected error with no tools configured: %v", errEvs)
}
}

View File

@ -1,595 +0,0 @@
package execute
import (
"context"
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
// TestMain points OLLIE_CFG_PATH at testdata so loadSandboxConfig uses the
// minimal test sandbox config rather than the user's ~/.config/ollie/sandbox/.
func TestMain(m *testing.M) {
wd, err := os.Getwd()
if err != nil {
panic(err)
}
os.Setenv("OLLIE_CFG_PATH", filepath.Join(wd, "testdata"))
os.Setenv("OLLIE", "")
os.Exit(m.Run())
}
// newServer returns a Server with Yolo=true so tests run without landrun in CI.
func newServer(t *testing.T) *Server {
t.Helper()
s := New(t.TempDir())
s.Yolo = true // skip landrun in unit tests
return s
}
// callCode invokes execute_code via Dispatch with a steps array.
func callCode(t *testing.T, s *Server, steps []map[string]any, extra ...map[string]any) (string, error) {
t.Helper()
payload := map[string]any{"steps": steps}
if len(extra) > 0 {
for k, v := range extra[0] {
payload[k] = v
}
}
raw, _ := json.Marshal(payload)
return s.Dispatch(context.Background(), "execute_code", raw)
}
// callTool invokes call_tool via Dispatch with a calls array.
func callTool(t *testing.T, s *Server, calls []map[string]any) (string, error) {
t.Helper()
payload := map[string]any{"calls": calls}
raw, _ := json.Marshal(payload)
return s.Dispatch(context.Background(), "call_tool", raw)
}
// callPipe invokes the pipe tool via Dispatch with a stages array.
func callPipe(t *testing.T, s *Server, stages []map[string]any) (string, error) {
t.Helper()
payload := map[string]any{"stages": stages}
raw, _ := json.Marshal(payload)
return s.Dispatch(context.Background(), "pipe", raw)
}
// ---- detectLanguage ----
func TestDetectLanguage(t *testing.T) {
cases := []struct{ code, want string }{
{"echo hi", "bash"},
{"#!/bin/bash\necho hi", "bash"},
{"#!/usr/bin/env python3\nprint(1)", "python3"},
{"#!/usr/bin/python\nprint(1)", "python3"},
{"#!/usr/bin/perl\nprint 1", "perl"},
{"#!/usr/bin/awk -f\n{print}", "awk"},
{"#!/usr/bin/env gawk\n{print}", "awk"},
{"#!/usr/bin/sed -f\ns/a/b/", "sed"},
{"#!/usr/bin/ed\n,p", "ed"},
{"#!/usr/bin/env jq\n.", "jq"},
{"#!/usr/bin/env lua\nprint(1)", "lua"},
}
for _, c := range cases {
if got := detectLanguage(c.code); got != c.want {
t.Errorf("detectLanguage(%q) = %q, want %q", c.code[:min(20, len(c.code))], got, c.want)
}
}
}
func min(a, b int) int {
if a < b {
return a
}
return b
}
// ---- ansiCEscape ----
func TestAnsiCEscape(t *testing.T) {
got := ansiCEscape("a\\b'c\nd\te")
want := `a\\b\'c\nd\te`
if got != want {
t.Errorf("got %q, want %q", got, want)
}
}
// ---- injectArgs ----
func TestInjectArgsBash(t *testing.T) {
out := injectArgs("bash", "myscript", []string{"hello", "world"}, "echo $1 $2")
if !strings.HasPrefix(out, "set -- ") {
t.Errorf("bash inject should start with 'set --', got: %q", out)
}
if !strings.Contains(out, "echo $1 $2") {
t.Errorf("bash inject should contain original code")
}
}
func TestInjectArgsPython(t *testing.T) {
out := injectArgs("python3", "s", []string{"a"}, "print(sys.argv)")
if !strings.HasPrefix(out, "import sys\n") {
t.Errorf("python inject should start with 'import sys', got: %q", out)
}
}
// ---- ToolsPath / PluginsPath ----
func TestToolsPathEnv(t *testing.T) {
t.Setenv("OLLIE_TOOLS_PATH", "/custom/tools:/other")
if got := ToolsPath(); got != "/custom/tools" {
t.Errorf("got %q, want /custom/tools", got)
}
}
func TestPluginsPathEnv(t *testing.T) {
t.Setenv("OLLIE_PLUGINS_PATH", "/custom/plugins:/other")
if got := PluginsPath(); got != "/custom/plugins" {
t.Errorf("got %q, want /custom/plugins", got)
}
}
// ---- ReadTool ----
func TestReadTool(t *testing.T) {
dir := t.TempDir()
t.Setenv("OLLIE_TOOLS_PATH", dir)
if err := os.WriteFile(filepath.Join(dir, "mytool"), []byte("#!/bin/bash\necho ok"), 0644); err != nil {
t.Fatal(err)
}
code, err := ReadTool("mytool")
if err != nil {
t.Fatal(err)
}
if !strings.Contains(code, "echo ok") {
t.Errorf("unexpected content: %q", code)
}
}
func TestReadToolNotFound(t *testing.T) {
t.Setenv("OLLIE_TOOLS_PATH", t.TempDir())
_, err := ReadTool("nonexistent")
if err == nil || !strings.Contains(err.Error(), "tool not found") {
t.Errorf("expected 'tool not found' error, got %v", err)
}
}
func TestReadToolInvalidName(t *testing.T) {
for _, name := range []string{"../etc/passwd", "foo/bar"} {
_, err := ReadTool(name)
if err == nil {
t.Errorf("expected error for name %q", name)
}
}
}
// ---- ValidateCode ----
func TestValidateCodeDangerous(t *testing.T) {
s := newServer(t)
cases := []struct{ code, lang string }{
{"sudo rm -rf /", "bash"},
{"rm -rf /home", "bash"},
{"mkfs /dev/sda", "bash"},
{"dd if=/dev/zero of=/dev/sda", "bash"},
{"shutil.rmtree('/')", "python3"},
}
for _, c := range cases {
if err := s.ValidateCode(c.code, c.lang); err == nil {
t.Errorf("expected dangerous pattern error for %q (%s)", c.code, c.lang)
}
}
}
func TestValidateCodeSafe(t *testing.T) {
s := newServer(t)
if err := s.ValidateCode("echo hello", "bash"); err != nil {
t.Errorf("unexpected error: %v", err)
}
}
// ---- Execute (integration, requires bash) ----
func TestExecuteSimpleBash(t *testing.T) {
s := newServer(t)
out, err := s.Execute(context.Background(), "echo hello", "bash", 10, "default", true)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if strings.TrimSpace(out) != "hello" {
t.Errorf("got %q, want %q", out, "hello")
}
}
func TestExecuteStdin(t *testing.T) {
s := newServer(t)
out, err := s.executeWithStdin(context.Background(), "cat", "bash", 10, "default", true, "piped input")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if strings.TrimSpace(out) != "piped input" {
t.Errorf("got %q, want %q", out, "piped input")
}
}
func TestExecuteTimeout(t *testing.T) {
s := newServer(t)
_, err := s.Execute(context.Background(), "sleep 10", "bash", 1, "default", true)
if err == nil || !strings.Contains(err.Error(), "timeout") {
t.Errorf("expected timeout error, got %v", err)
}
}
func TestExecuteNonZeroExit(t *testing.T) {
s := newServer(t)
_, err := s.Execute(context.Background(), "exit 1", "bash", 10, "default", true)
if err == nil {
t.Error("expected error for non-zero exit")
}
}
func TestExecuteUnsupportedLanguage(t *testing.T) {
s := newServer(t)
_, err := s.Execute(context.Background(), "code", "cobol", 10, "default", true)
if err == nil || !strings.Contains(err.Error(), "unsupported language") {
t.Errorf("expected unsupported language error, got %v", err)
}
}
// ---- Dispatch / execute_code ----
func TestDispatchUnknownTool(t *testing.T) {
s := newServer(t)
_, err := s.Dispatch(context.Background(), "unknown_tool", json.RawMessage(`{}`))
if err == nil {
t.Error("expected error for unknown tool")
}
}
func TestDispatchNoSteps(t *testing.T) {
s := newServer(t)
_, err := callCode(t, s, []map[string]any{})
if err == nil || !strings.Contains(err.Error(), "steps is required") {
t.Errorf("expected 'steps is required' error, got %v", err)
}
}
func TestDispatchSingleStep(t *testing.T) {
s := newServer(t)
out, err := callCode(t, s, []map[string]any{{"code": "echo single"}})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !strings.Contains(out, "single") {
t.Errorf("got %q, want output containing 'single'", out)
}
}
// ---- Dispatch / pipe ----
func TestDispatchPipeline(t *testing.T) {
s := newServer(t)
out, err := callPipe(t, s, []map[string]any{
{"code": "printf 'a\\nb\\nc'"},
{"code": "grep b"},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if strings.TrimSpace(out) != "b" {
t.Errorf("got %q, want %q", out, "b")
}
}
func TestDispatchPipeNoStages(t *testing.T) {
s := newServer(t)
_, err := callPipe(t, s, []map[string]any{})
if err == nil || !strings.Contains(err.Error(), "stages is required") {
t.Errorf("expected 'stages is required' error, got %v", err)
}
}
// ---- Dispatch / execute_code parallel ----
func TestDispatchParallel(t *testing.T) {
s := newServer(t)
out, err := callCode(t, s, []map[string]any{
{"parallel": []map[string]any{
{"code": "echo A"},
{"code": "echo B"},
}},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !strings.Contains(out, "A") || !strings.Contains(out, "B") {
t.Errorf("parallel output missing A or B: %q", out)
}
}
// ---- Dispatch / call_tool ----
func TestDispatchCallTool(t *testing.T) {
dir := t.TempDir()
t.Setenv("OLLIE_TOOLS_PATH", dir)
if err := os.WriteFile(filepath.Join(dir, "greet"), []byte("#!/bin/bash\necho hello-from-tool"), 0755); err != nil {
t.Fatal(err)
}
s := newServer(t)
out, err := callTool(t, s, []map[string]any{{"tool": "greet"}})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !strings.Contains(out, "hello-from-tool") {
t.Errorf("got %q, want output containing 'hello-from-tool'", out)
}
}
func TestDispatchCallToolNoCalls(t *testing.T) {
s := newServer(t)
_, err := callTool(t, s, []map[string]any{})
if err == nil || !strings.Contains(err.Error(), "calls is required") {
t.Errorf("expected 'calls is required' error, got %v", err)
}
}
func TestDispatchCallToolInlineRejected(t *testing.T) {
s := newServer(t)
_, err := callTool(t, s, []map[string]any{{"code": "echo hi"}})
if err == nil || !strings.Contains(err.Error(), "tool name is required") {
t.Errorf("expected 'tool name is required' error for inline code in call_tool, got %v", err)
}
}
// The old TestDispatchToolStep tested tool steps via execute_code.
// That path still works (execute_code accepts tool steps in its steps array);
// but call_tool is now the canonical way. Test both for coverage.
func TestDispatchToolStepViaExecuteCode(t *testing.T) {
dir := t.TempDir()
t.Setenv("OLLIE_TOOLS_PATH", dir)
if err := os.WriteFile(filepath.Join(dir, "greet"), []byte("#!/bin/bash\necho hello-from-tool"), 0755); err != nil {
t.Fatal(err)
}
s := newServer(t)
out, err := callCode(t, s, []map[string]any{{"tool": "greet"}})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !strings.Contains(out, "hello-from-tool") {
t.Errorf("got %q, want output containing 'hello-from-tool'", out)
}
}
// ---- limitedWriter ----
func TestLimitedWriter(t *testing.T) {
var buf strings.Builder
lw := &limitedWriter{w: &buf, limit: 5}
lw.Write([]byte("hello world"))
if buf.String() != "hello" {
t.Errorf("got %q, want %q", buf.String(), "hello")
}
if !lw.truncated {
t.Error("expected truncated=true")
}
}
// ---- rate limiting ----
func TestRateLimitBlocks(t *testing.T) {
s := newServer(t)
// Trigger maxFailures validation failures.
for i := 0; i < maxFailures; i++ {
s.recordValidationFailure()
}
if err := s.checkRateLimit(); err == nil {
t.Error("expected rate limit error after max failures")
}
}
// ---- SetEnv / SetCWD ----
func TestSetEnvInjected(t *testing.T) {
s := newServer(t)
s.SetEnv("MY_TEST_VAR", "injected_value")
out, err := s.Execute(context.Background(), "echo $MY_TEST_VAR", "bash", 10, "default", true)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !strings.Contains(out, "injected_value") {
t.Errorf("got %q, want 'injected_value'", out)
}
}
func TestSetCWD(t *testing.T) {
s := newServer(t)
dir := t.TempDir()
s.SetCWD(dir)
out, err := s.Execute(context.Background(), "pwd", "bash", 10, "default", true)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
// TempDir may use a symlink; compare base name only.
if !strings.Contains(out, filepath.Base(dir)) {
t.Errorf("got %q, want path containing %q", out, filepath.Base(dir))
}
}
// ---- Strict mode ----
func TestStrictModeRejectsInline(t *testing.T) {
s := newServer(t)
s.Strict = true
_, err := callCode(t, s, []map[string]any{{"code": "echo hi"}})
if err == nil || !strings.Contains(err.Error(), "strict") {
t.Errorf("expected strict mode error, got %v", err)
}
}
func TestStrictModeAllowsTool(t *testing.T) {
dir := t.TempDir()
t.Setenv("OLLIE_TOOLS_PATH", dir)
if err := os.WriteFile(filepath.Join(dir, "safe"), []byte("#!/bin/bash\necho safe"), 0755); err != nil {
t.Fatal(err)
}
s := newServer(t)
s.Strict = true
out, err := callCode(t, s, []map[string]any{{"tool": "safe"}})
if err != nil {
t.Fatalf("strict mode should allow tool steps: %v", err)
}
if !strings.Contains(out, "safe") {
t.Errorf("got %q, want 'safe'", out)
}
}
// ---- classifyStep ----
func TestClassifyStep(t *testing.T) {
// inline code is unclassifiable without scanning → global
if got := classifyStep(CodeStep{Code: "echo hi"}); got != lockClassGlobal {
t.Errorf("inline echo: got %v; want global", got)
}
// elevated → global
if got := classifyStep(CodeStep{Elevated: true}); got != lockClassGlobal {
t.Errorf("elevated: got %v; want global", got)
}
// unknown tool → global
if got := classifyStep(CodeStep{Tool: "nonexistent_tool_xyz"}); got != lockClassGlobal {
t.Errorf("unknown tool: got %v; want global", got)
}
}
func TestSetLockDir(t *testing.T) {
s := newServer(t)
dir := t.TempDir()
s.SetLockDir(dir)
if s.lockDir != dir {
t.Errorf("lockDir = %q; want %q", s.lockDir, dir)
}
}
func TestRunReadBatch_Single(t *testing.T) {
s := newServer(t)
ctx := context.Background()
stages := []CodeStep{{Code: "echo batch", Tool: ""}}
out, err := s.runReadBatch(ctx, 0, stages, 10, "default", "")
if err != nil {
t.Fatalf("runReadBatch: %v", err)
}
if !strings.Contains(out, "batch") {
t.Errorf("output = %q; want 'batch'", out)
}
}
func TestRunReadBatch_Multi(t *testing.T) {
s := newServer(t)
ctx := context.Background()
stages := []CodeStep{
{Code: "echo one"},
{Code: "echo two"},
}
out, err := s.runReadBatch(ctx, 0, stages, 10, "default", "")
if err != nil {
t.Fatalf("runReadBatch multi: %v", err)
}
if !strings.Contains(out, "one") || !strings.Contains(out, "two") {
t.Errorf("output = %q; want 'one' and 'two'", out)
}
}
func TestRunReadBatch_WithLockDir(t *testing.T) {
s := newServer(t)
s.SetLockDir(t.TempDir())
ctx := context.Background()
stages := []CodeStep{{Code: "echo locked"}}
out, err := s.runReadBatch(ctx, 0, stages, 10, "default", "")
if err != nil {
t.Fatalf("runReadBatch with lockdir: %v", err)
}
if !strings.Contains(out, "locked") {
t.Errorf("output = %q; want 'locked'", out)
}
}
func TestTimeoutZeroNoDeadline(t *testing.T) {
s := newServer(t)
ctx := context.Background()
// timeout=0 should not kill the process after 30s default.
// We verify by running a 2s sleep with timeout=0 and detaching it.
// If timeout were applied, it would use default 30s which is fine,
// but we confirm it doesn't fail with timeout=0.
done := make(chan string, 1)
go func() {
out, _ := s.Execute(ctx, "echo alive; sleep 1; echo done", "bash", 0, "default", true)
done <- out
}()
out := <-done
if !strings.Contains(out, "done") {
t.Errorf("expected 'done' in output, got: %q", out)
}
}
// ---- detach ----
func TestDetach(t *testing.T) {
s := newServer(t)
ctx := context.Background()
// Start a long-running process, then detach it
var detachedPID int
s.OnDetach = func(pid int, cmd string) {
detachedPID = pid
}
done := make(chan string, 1)
go func() {
out, _ := s.Execute(ctx, "echo started; sleep 30", "bash", 60, "default", true)
done <- out
}()
// Give the process time to start and reach the select
for i := 0; i < 50; i++ {
if s.Detach() {
break
}
time.Sleep(10 * time.Millisecond)
}
out := <-done
if !strings.Contains(out, "[detached: pid") {
t.Fatalf("expected detach marker in output, got: %q", out)
}
if detachedPID == 0 {
t.Fatal("OnDetach not called")
}
// Verify process is in detached list
procs := s.ListDetached()
if len(procs) != 1 {
t.Fatalf("expected 1 detached process, got %d", len(procs))
}
if procs[0].PID != detachedPID {
t.Errorf("PID mismatch: got %d, want %d", procs[0].PID, detachedPID)
}
if procs[0].Exited {
t.Error("process should still be running")
}
// Kill it
if err := s.SignalDetached(detachedPID, 15); err != nil { // SIGTERM
t.Fatalf("SignalDetached: %v", err)
}
// Wait for exit
<-procs[0].done
if !procs[0].Exited {
t.Error("process should have exited after SIGTERM")
}
}

View File

@ -1,56 +0,0 @@
package execute
import (
"os"
"path/filepath"
"strings"
"syscall"
)
// acquireFlock opens (or creates) a lock file in dir named name and acquires
// LOCK_SH (exclusive=false) or LOCK_EX (exclusive=true).
// Returns nil, nil when dir is empty (locking disabled).
// Caller must Close the returned file to release the lock.
func acquireFlock(dir, name string, exclusive bool) (*os.File, error) {
if dir == "" {
return nil, nil
}
if err := os.MkdirAll(dir, 0700); err != nil {
return nil, err
}
how := syscall.LOCK_SH
if exclusive {
how = syscall.LOCK_EX
}
path := filepath.Join(dir, sanitizeLockName(name)+".lock")
f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0600)
if err != nil {
return nil, err
}
if err := syscall.Flock(int(f.Fd()), how); err != nil {
f.Close()
return nil, err
}
return f, nil
}
func sanitizeLockName(s string) string {
var b strings.Builder
for _, r := range s {
switch r {
case '/', '\\', ':', '*', '?', '<', '>', '|', '"', ' ':
b.WriteRune('_')
default:
b.WriteRune(r)
}
}
name := b.String()
if len(name) > 64 {
name = name[:64]
}
if name == "" {
return "unnamed"
}
return name
}

View File

@ -1,221 +0,0 @@
package execute
import (
"os"
"testing"
)
func init() {
os.Setenv("OLLIE_TOOLS_PATH", "testdata/tools")
}
func TestOrderArgsBySchema(t *testing.T) {
script := `#!/usr/bin/env bash
# args_json: {"type":"object","properties":{"pattern":{"type":"string"},"path":{"type":"string"}},"required":["pattern"]}
# ollie:prompt
# ## test_tool
# ollie:end
`
// Run 100 times to ensure deterministic ordering
for i := 0; i < 100; i++ {
argMap := map[string]interface{}{
"pattern": "*.go",
"path": "/some/dir",
}
result := orderArgsBySchema(script, argMap)
if len(result) != 2 {
t.Fatalf("iteration %d: expected 2 args, got %d: %v", i, len(result), result)
}
if result[0] != "*.go" {
t.Fatalf("iteration %d: expected result[0]='*.go', got %q", i, result[0])
}
if result[1] != "/some/dir" {
t.Fatalf("iteration %d: expected result[1]='/some/dir', got %q", i, result[1])
}
}
}
func TestOrderArgsBySchemaReversed(t *testing.T) {
script := `#!/usr/bin/env bash
# args_json: {"type":"object","properties":{"path":{"type":"string"},"pattern":{"type":"string"}},"required":["path"]}
# ollie:prompt
# ## test_tool
# ollie:end
`
for i := 0; i < 100; i++ {
argMap := map[string]interface{}{
"pattern": "*.go",
"path": "/some/dir",
}
result := orderArgsBySchema(script, argMap)
if len(result) != 2 {
t.Fatalf("iteration %d: expected 2 args, got %d: %v", i, len(result), result)
}
if result[0] != "/some/dir" {
t.Fatalf("iteration %d: expected result[0]='/some/dir', got %q", i, result[0])
}
if result[1] != "*.go" {
t.Fatalf("iteration %d: expected result[1]='*.go', got %q", i, result[1])
}
}
}
func TestRegistryLoadUnload(t *testing.T) {
r, err := NewRegistry()
if err != nil {
t.Fatalf("NewRegistry: %v", err)
}
// Check global discovery
names := r.GlobalToolNames()
if len(names) == 0 {
t.Fatal("expected at least one tool")
}
// Load a tool
sid := "test-session-1"
err = r.Load(sid, names[0])
if err != nil {
t.Fatalf("Load: %v", err)
}
// Loaded should have 1 tool
loaded := r.Loaded(sid)
if len(loaded) != 1 {
t.Fatalf("expected 1 loaded tool, got %d", len(loaded))
}
if loaded[0].Name != names[0] {
t.Fatalf("expected %s, got %s", names[0], loaded[0].Name)
}
// Revision should be 1
rev := r.Revision(sid)
if rev != 1 {
t.Fatalf("expected revision 1, got %d", rev)
}
// Lookup should work
_, ok := r.Lookup(sid, names[0])
if !ok {
t.Fatal("Lookup should find promoted tool")
}
// Unload
err = r.Unload(sid, names[0])
if err != nil {
t.Fatalf("Unload: %v", err)
}
// Revision should be 2
rev = r.Revision(sid)
if rev != 2 {
t.Fatalf("expected revision 2, got %d", rev)
}
// Loaded should be empty
loaded = r.Loaded(sid)
if len(loaded) != 0 {
t.Fatalf("expected 0 loaded tools, got %d", len(loaded))
}
}
func TestRegistrySummaries(t *testing.T) {
r, err := NewRegistry()
if err != nil {
t.Fatalf("NewRegistry: %v", err)
}
summaries := r.Summaries()
if len(summaries) == 0 {
t.Fatal("expected at least one summary")
}
for _, s := range summaries {
if s.Name == "" {
t.Fatal("summary missing name")
}
}
}
func TestRegistryIdempotentLoad(t *testing.T) {
r, err := NewRegistry()
if err != nil {
t.Fatalf("NewRegistry: %v", err)
}
names := r.GlobalToolNames()
if len(names) == 0 {
t.Fatal("expected at least one tool")
}
sid := "test-session-2"
// Load twice
err = r.Load(sid, names[0])
if err != nil {
t.Fatalf("first Load: %v", err)
}
rev1 := r.Revision(sid)
err = r.Load(sid, names[0])
if err != nil {
t.Fatalf("second Load: %v", err)
}
rev2 := r.Revision(sid)
if rev1 != rev2 {
t.Fatal("idempotent load should not bump revision")
}
}
func TestRegistryToolNotFound(t *testing.T) {
r, err := NewRegistry()
if err != nil {
t.Fatalf("NewRegistry: %v", err)
}
err = r.Load("test", "nonexistent-tool")
if err == nil {
t.Fatal("expected error for nonexistent tool")
}
}
func TestRegistryServerListTools(t *testing.T) {
r, err := NewRegistry()
if err != nil {
t.Fatalf("NewRegistry: %v", err)
}
names := r.GlobalToolNames()
if len(names) == 0 {
t.Fatal("expected at least one tool")
}
sid := "test-session-3"
r.Load(sid, names[0])
srv := &Server{}
srv.SetToolRegistry(r, sid)
tools, err := srv.ListTools()
if err != nil {
t.Fatalf("ListTools: %v", err)
}
if len(tools) < 1 {
t.Fatalf("expected at least 1 tool, got %d", len(tools))
}
// First tool is execute_code, promoted tools come after
found := false
for _, ti := range tools {
if ti.Name == names[0] {
found = true
break
}
}
if !found {
t.Fatalf("expected promoted tool %s in list", names[0])
}
if tools[0].InputSchema == nil {
t.Fatal("expected InputSchema to be non-nil")
}
}

View File

@ -44,6 +44,25 @@ func ExtractReturnSchema(script string) json.RawMessage {
return nil
}
func ExtractTier(script string) string {
for _, line := range strings.Split(script, "\n") {
idx := strings.Index(line, "ollie:tier")
if idx < 0 {
continue
}
rest := strings.TrimSpace(line[idx+len("ollie:tier"):])
switch {
case rest == "cold" || strings.HasPrefix(rest, "cold "):
return "cold"
case rest == "warm" || strings.HasPrefix(rest, "warm "):
return "warm"
case rest == "hot" || strings.HasPrefix(rest, "hot "):
return "hot"
}
}
return ""
}
func ExtractMetadata(script string) ToolMeta {
var meta ToolMeta
for _, line := range strings.Split(script, "\n") {
@ -62,10 +81,17 @@ func ParseToolInfo(name, script string) tools.ToolInfo {
if argsSchema == nil {
argsSchema = json.RawMessage(`{"type":"object","properties":{"tool":{"type":"string"},"args":{"type":"array","items":{"type":"string"}}},"required":["tool","args"]}`)
}
meta := ExtractMetadata(script)
tier := ExtractTier(script)
if tier == "" {
tier = "hot"
}
return tools.ToolInfo{
Name: name,
Description: desc,
InputSchema: argsSchema,
Prompt: prompt,
Tier: tier,
ReadOnly: meta.ReadOnly,
}
}

View File

@ -39,10 +39,6 @@ type Server struct {
envMu sync.RWMutex
envExtra map[string]string
// lockDir is the directory for advisory flock files used during parallel
// step dispatch. Set once at session init via SetLockDir; empty disables locking.
lockDir string
// Hooks for lifecycle events (harnesses like 9P can inject mount logic here)
OnPreDispatch func()
OnClose func()
@ -210,10 +206,6 @@ func (e *Server) SetAllowTools(names []string) {
}
}
// SetLockDir sets the directory used for advisory flock files during parallel
// step dispatch. Must be called before the Server handles concurrent requests.
func (e *Server) SetLockDir(dir string) { e.lockDir = dir }
// SetEnv adds a session-scoped environment variable injected into all
// subsequent subprocess invocations for this session.
func (e *Server) SetEnv(key, value string) {

View File

@ -5,29 +5,13 @@ import (
"strings"
)
// defaultColdTools are tools whose results are consumed immediately and don't
// need to persist verbatim in the message history.
var defaultColdTools = map[string]string{
"file_read": "cold",
"file_grep": "cold",
"file_glob": "cold",
"memory_recall": "cold",
"web_fetch": "cold",
"web_search": "cold",
"lsp_definition": "cold",
"lsp_references": "cold",
"lsp_hover": "cold",
"lsp_symbols": "cold",
"lsp_completion": "cold",
"lsp_diagnostics": "cold",
"reasoning_think": "cold",
}
// ResultTier implements tools.TierClassifier. Checks the built-in table first,
// then falls back to the script's "ollie:tier" annotation.
// ResultTier implements tools.TierClassifier. Looks up the tool's cached tier
// from the registry; falls back to reading the script on disk.
func (e *Server) ResultTier(name string) string {
if tier, ok := defaultColdTools[name]; ok {
return tier
if e.toolRegistry != nil && e.sessionID != "" {
if info, ok := e.toolRegistry.Lookup(e.sessionID, name); ok && info.Tier != "" {
return info.Tier
}
}
code, err := ReadTool(name)
if err != nil {
@ -48,24 +32,31 @@ func (e *Server) ResultTierArgs(name string, args json.RawMessage) string {
}
}
// colderTier returns the colder of two tiers. cold < warm < hot.
func colderTier(a, b string) string {
order := map[string]int{"cold": 0, "warm": 1, "hot": 2}
if order[b] < order[a] {
return b
// IsParallelRead implements tools.ParallelClassifier. Returns true when the
// named tool script is ReadOnly (ollie:parallel read annotation).
// Checks the registry cache first; falls back to reading the script.
func (e *Server) IsParallelRead(name string) bool {
if e.toolRegistry != nil && e.sessionID != "" {
if info, ok := e.toolRegistry.Lookup(e.sessionID, name); ok {
return info.ReadOnly
}
}
return a
code, err := ReadTool(name)
if err != nil {
return false
}
for _, line := range strings.SplitN(code, "\n", 11) {
if strings.Contains(line, "ollie:parallel read") {
return true
}
}
return false
}
// detectTier scans the first 10 lines of a tool script for an
// "ollie:tier" annotation. Matches "ollie:tier hot", "ollie:tier warm",
// or "ollie:tier cold" anywhere in the line.
// "ollie:tier" annotation. Fallback when a tool isn't in the registry.
func detectTier(code string) string {
lines := strings.SplitN(code, "\n", 11)
if len(lines) > 10 {
lines = lines[:10]
}
for _, line := range lines {
for _, line := range strings.SplitN(code, "\n", 11) {
idx := strings.Index(line, "ollie:tier")
if idx < 0 {
continue

View File

@ -1,97 +0,0 @@
package execute
import (
"encoding/json"
"testing"
)
func TestResultTierArgs_CallToolCold(t *testing.T) {
s := &Server{}
args := json.RawMessage(`{"calls":[{"tool":"file_read","args":["/tmp/x"]}]}`)
if got := s.ResultTierArgs("call_tool", args); got != "cold" {
t.Errorf("ResultTierArgs(call_tool, file_read) = %q; want cold", got)
}
}
func TestResultTierArgs_CallToolMultipleCold(t *testing.T) {
s := &Server{}
args := json.RawMessage(`{"calls":[{"tool":"file_read","args":["/a"]},{"tool":"file_grep","args":["pat"]}]}`)
if got := s.ResultTierArgs("call_tool", args); got != "cold" {
t.Errorf("ResultTierArgs(call_tool, file_read+file_grep) = %q; want cold", got)
}
}
func TestResultTierArgs_CallToolMixedHot(t *testing.T) {
s := &Server{}
// file_write is not in defaultColdTools and has no script, so falls back to hot
args := json.RawMessage(`{"calls":[{"tool":"file_read","args":["/a"]},{"tool":"file_write","args":["/b","x"]}]}`)
if got := s.ResultTierArgs("call_tool", args); got != "cold" {
t.Errorf("ResultTierArgs(call_tool, file_read+file_write) = %q; want cold (coldest wins)", got)
}
}
func TestResultTierArgs_CallToolParallel(t *testing.T) {
s := &Server{}
args := json.RawMessage(`{"calls":[{"parallel":[{"tool":"file_grep","args":["x"]},{"tool":"lsp_hover","args":["/f","1","1"]}]}]}`)
if got := s.ResultTierArgs("call_tool", args); got != "cold" {
t.Errorf("ResultTierArgs(call_tool, parallel cold tools) = %q; want cold", got)
}
}
func TestResultTierArgs_PipeCold(t *testing.T) {
s := &Server{}
args := json.RawMessage(`{"stages":[{"tool":"file_grep","args":["TODO"]},{"code":"wc -l"}]}`)
if got := s.ResultTierArgs("pipe", args); got != "cold" {
t.Errorf("ResultTierArgs(pipe, file_grep+code) = %q; want cold", got)
}
}
func TestResultTierArgs_ExecuteCode(t *testing.T) {
s := &Server{}
args := json.RawMessage(`{"steps":[{"code":"date"}]}`)
if got := s.ResultTierArgs("execute_code", args); got != "warm" {
t.Errorf("ResultTierArgs(execute_code) = %q; want warm", got)
}
}
func TestResultTierArgs_BadJSON(t *testing.T) {
s := &Server{}
if got := s.ResultTierArgs("call_tool", json.RawMessage(`{invalid`)); got != "hot" {
t.Errorf("ResultTierArgs(bad json) = %q; want hot", got)
}
}
func TestResultTierArgs_EmptyCalls(t *testing.T) {
s := &Server{}
if got := s.ResultTierArgs("call_tool", json.RawMessage(`{"calls":[]}`)); got != "hot" {
t.Errorf("ResultTierArgs(empty calls) = %q; want hot", got)
}
}
func TestResultTier_DefaultColdTools(t *testing.T) {
s := &Server{}
coldTools := []string{"file_read", "file_grep", "file_glob", "memory_recall",
"web_fetch", "web_search", "lsp_definition", "lsp_references",
"lsp_hover", "lsp_symbols", "lsp_completion", "lsp_diagnostics", "reasoning_think"}
for _, name := range coldTools {
if got := s.ResultTier(name); got != "cold" {
t.Errorf("ResultTier(%q) = %q; want cold", name, got)
}
}
}
func TestColderTier(t *testing.T) {
tests := []struct{ a, b, want string }{
{"hot", "cold", "cold"},
{"cold", "hot", "cold"},
{"warm", "cold", "cold"},
{"hot", "warm", "warm"},
{"hot", "hot", "hot"},
{"cold", "cold", "cold"},
}
for _, tt := range tests {
if got := colderTier(tt.a, tt.b); got != tt.want {
t.Errorf("colderTier(%q, %q) = %q; want %q", tt.a, tt.b, got, tt.want)
}
}
}

View File

@ -23,36 +23,6 @@ func ToolsPath() string {
return paths.CfgDir() + "/tools"
}
// PluginsPath returns the directory for server-invoked plugin scripts.
// Resolved in order: first entry of OLLIE_PLUGINS_PATH (colon-separated),
// then ~/.config/ollie/scripts/x.
func PluginsPath() string {
if p := os.Getenv("OLLIE_PLUGINS_PATH"); p != "" {
if i := strings.Index(p, ":"); i >= 0 {
p = p[:i]
}
return p
}
return paths.CfgDir() + "/scripts/x"
}
// injectArgs prepends argument binding to code.
func injectArgs(language, name string, args []string, code string) string {
escaped := make([]string, len(args))
for i, a := range args {
escaped[i] = "'" + strings.ReplaceAll(a, "'", "'\\''") + "'"
}
return fmt.Sprintf("set -- %s\n%s", strings.Join(escaped, " "), code)
}
// InjectArgs is the exported version of injectArgs.
func InjectArgs(language, name string, args []string, code string) string {
return injectArgs(language, name, args, code)
}
// ExtractPrompt parses the ollie:prompt ... ollie:end block from a tool
// script's header comments. Returns the prompt text with comment prefixes

View File

@ -17,6 +17,11 @@ type ToolInfo struct {
// Prompt is the usage documentation for this tool, extracted from
// the script's ollie:prompt block. Included in the system prompt.
Prompt string
// Tier is the retention tier: "hot", "warm", or "cold". Parsed from
// the script's ollie:tier annotation. Empty defaults to "hot".
Tier string
// ReadOnly is true when the tool carries an "ollie:parallel read" annotation.
ReadOnly bool
}
// Server is the interface satisfied by any tool server.
@ -104,12 +109,6 @@ type EnvSetter interface {
SetEnv(key, value string)
}
// LockDirSetter is implemented by tool servers that use advisory flock files
// for parallel-step coordination. SetLockDir must be called once at session init.
type LockDirSetter interface {
SetLockDir(dir string)
}
// ToolRestrictionSetter is implemented by tool servers that support restricting
// which tool scripts are available.
type ToolRestrictionSetter interface {