tool registry: cache tier/parallel metadata in ToolInfo, drop hard-coded tables
This commit is contained in:
parent
2c36c023aa
commit
ba9f605119
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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")
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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")
|
||||
}
|
||||
}
|
||||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
Reference in New Issue