merge registry into execute: one server, one routing table, deterministic arg ordering
This commit is contained in:
parent
f08e0abdf0
commit
dce77f0b8d
|
|
@ -1,4 +1,4 @@
|
|||
package registry
|
||||
package execute
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
|
@ -7,7 +7,7 @@ import (
|
|||
"sync"
|
||||
|
||||
"ollie/pkg/tools"
|
||||
"ollie/pkg/tools/execute"
|
||||
|
||||
)
|
||||
|
||||
type Registry struct {
|
||||
|
|
@ -30,7 +30,7 @@ func NewRegistry() (*Registry, error) {
|
|||
}
|
||||
|
||||
func (r *Registry) Discover() error {
|
||||
dir := execute.ToolsPath()
|
||||
dir := ToolsPath()
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read tools dir %s: %w", dir, err)
|
||||
|
|
@ -1,9 +1,65 @@
|
|||
package registry
|
||||
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 {
|
||||
|
|
@ -138,20 +194,26 @@ func TestRegistryServerListTools(t *testing.T) {
|
|||
sid := "test-session-3"
|
||||
r.Load(sid, names[0])
|
||||
|
||||
srv := &Server{
|
||||
registry: r,
|
||||
sessionID: sid,
|
||||
}
|
||||
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 1 tool, got %d", len(tools))
|
||||
if len(tools) < 1 {
|
||||
t.Fatalf("expected at least 1 tool, got %d", len(tools))
|
||||
}
|
||||
if tools[0].Name != names[0] {
|
||||
t.Fatalf("expected %s, got %s", names[0], tools[0].Name)
|
||||
// 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")
|
||||
|
|
@ -1,11 +1,11 @@
|
|||
package registry
|
||||
package execute
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"ollie/pkg/tools"
|
||||
"ollie/pkg/tools/execute"
|
||||
|
||||
)
|
||||
|
||||
type ToolMeta struct {
|
||||
|
|
@ -56,8 +56,8 @@ func ExtractMetadata(script string) ToolMeta {
|
|||
}
|
||||
|
||||
func ParseToolInfo(name, script string) tools.ToolInfo {
|
||||
prompt := execute.ExtractPrompt(script)
|
||||
desc := execute.ExtractShortDescription(prompt)
|
||||
prompt := ExtractPrompt(script)
|
||||
desc := ExtractShortDescription(prompt)
|
||||
argsSchema := ExtractArgsSchema(script)
|
||||
if argsSchema == nil {
|
||||
argsSchema = json.RawMessage(`{"type":"object","properties":{"tool":{"type":"string"},"args":{"type":"array","items":{"type":"string"}}},"required":["tool","args"]}`)
|
||||
|
|
@ -1,9 +1,9 @@
|
|||
package execute
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
|
|
@ -62,6 +62,10 @@ type Server struct {
|
|||
// Empty means all are allowed.
|
||||
allowTools map[string]bool
|
||||
|
||||
toolRegistry *Registry
|
||||
|
||||
sessionID string
|
||||
|
||||
// rate limiting state (per-Server)
|
||||
rateLimitMu sync.Mutex
|
||||
validationFailures int
|
||||
|
|
@ -144,8 +148,8 @@ func Decl(cwd string, opts ...Option) func() tools.Server {
|
|||
}
|
||||
}
|
||||
|
||||
// ListTools implements tools.Server, returning ToolInfo for execute_code,
|
||||
// call_tool, and pipe.
|
||||
// ListTools implements tools.Server, returning execute_code plus any
|
||||
// tools promoted in the session's tool registry.
|
||||
func (e *Server) ListTools() ([]tools.ToolInfo, error) {
|
||||
all := []tools.ToolInfo{
|
||||
{
|
||||
|
|
@ -154,8 +158,6 @@ func (e *Server) ListTools() ([]tools.ToolInfo, error) {
|
|||
|
||||
Steps run in parallel when consecutive steps carry an ollie:parallel annotation;
|
||||
otherwise they run serially. Outputs are concatenated in submission order.
|
||||
No stdout chaining between steps — use pipe for that.
|
||||
|
||||
Each step is one of:
|
||||
- {code, language} — inline code (default language: bash)
|
||||
- {elevated: true, code} — run outside sandbox via elevation backend (bash only)
|
||||
|
|
@ -323,11 +325,19 @@ Examples:
|
|||
}
|
||||
return filtered, nil
|
||||
}
|
||||
if e.toolRegistry != nil && e.sessionID != "" {
|
||||
all = append(all, e.toolRegistry.Loaded(e.sessionID)...)
|
||||
}
|
||||
return all, nil
|
||||
}
|
||||
|
||||
// CallTool implements tools.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 {
|
||||
return e.callPromotedTool(ctx, tool, args)
|
||||
}
|
||||
}
|
||||
result, err := e.Dispatch(ctx, tool, args)
|
||||
if err != nil {
|
||||
return json.Marshal(map[string]any{
|
||||
|
|
@ -399,6 +409,100 @@ func (e *Server) SetEnv(key, value string) {
|
|||
}
|
||||
}
|
||||
|
||||
// SetToolRegistry attaches a session-local tool registry.
|
||||
func (e *Server) SetToolRegistry(r *Registry, sessionID string) {
|
||||
e.toolRegistry = r
|
||||
e.sessionID = sessionID
|
||||
}
|
||||
|
||||
// callPromotedTool executes a tool promoted via the registry.
|
||||
func (e *Server) callPromotedTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) {
|
||||
script, err := ReadTool(tool)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read tool %s: %w", tool, err)
|
||||
}
|
||||
var positional []string
|
||||
var argMap map[string]interface{}
|
||||
if err := json.Unmarshal(args, &argMap); err == nil {
|
||||
if argsArr, ok := argMap["args"]; ok {
|
||||
if arr, ok := argsArr.([]interface{}); ok {
|
||||
for _, v := range arr {
|
||||
positional = append(positional, fmt.Sprintf("%v", v))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
positional = orderArgsBySchema(script, argMap)
|
||||
}
|
||||
}
|
||||
lang := DetectLanguage(script)
|
||||
code := InjectArgs(lang, tool, positional, script)
|
||||
result, err := e.Execute(ctx, code, lang, 30, "default", true)
|
||||
if err != nil {
|
||||
return json.Marshal(map[string]interface{}{
|
||||
"isError": true,
|
||||
"content": []map[string]string{{"type": "text", "text": err.Error()}},
|
||||
})
|
||||
}
|
||||
return json.Marshal(map[string]interface{}{
|
||||
"content": []map[string]string{{"type": "text", "text": result}},
|
||||
})
|
||||
}
|
||||
|
||||
// orderArgsBySchema extracts values from argMap in schema property order.
|
||||
// Uses json.Decoder to preserve JSON key ordering.
|
||||
func orderArgsBySchema(script string, argMap map[string]interface{}) []string {
|
||||
schema := ExtractArgsSchema(script)
|
||||
if schema == nil {
|
||||
var out []string
|
||||
for _, v := range argMap {
|
||||
out = append(out, fmt.Sprintf("%v", v))
|
||||
}
|
||||
return out
|
||||
}
|
||||
// Parse the schema to find the "properties" key
|
||||
var raw map[string]json.RawMessage
|
||||
if err := json.Unmarshal(schema, &raw); err != nil {
|
||||
var out []string
|
||||
for _, v := range argMap {
|
||||
out = append(out, fmt.Sprintf("%v", v))
|
||||
}
|
||||
return out
|
||||
}
|
||||
props, ok := raw["properties"]
|
||||
if !ok || len(props) == 0 {
|
||||
var out []string
|
||||
for _, v := range argMap {
|
||||
out = append(out, fmt.Sprintf("%v", v))
|
||||
}
|
||||
return out
|
||||
}
|
||||
// Decode properties sub-object preserving key order using json.Decoder
|
||||
dec := json.NewDecoder(bytes.NewReader(props))
|
||||
tok, err := dec.Token()
|
||||
if err != nil || tok != json.Delim('{') {
|
||||
var out []string
|
||||
for _, v := range argMap {
|
||||
out = append(out, fmt.Sprintf("%v", v))
|
||||
}
|
||||
return out
|
||||
}
|
||||
var out []string
|
||||
for dec.More() {
|
||||
keyTok, err := dec.Token()
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
key := fmt.Sprintf("%v", keyTok)
|
||||
// Skip the value
|
||||
var val json.RawMessage
|
||||
dec.Decode(&val)
|
||||
if v, ok := argMap[key]; ok {
|
||||
out = append(out, fmt.Sprintf("%v", v))
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// Close is called when the session ends. Calls OnClose hook if registered.
|
||||
func (e *Server) Close() {
|
||||
e.cleanupDetached()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,8 @@
|
|||
#!/usr/bin/env bash
|
||||
# args_json: {"type":"object","properties":{"input":{"type":"string"}},"required":["input"]}
|
||||
# ollie:prompt
|
||||
# ## test_tool
|
||||
#
|
||||
# A test tool for registry tests.
|
||||
# ollie:end
|
||||
echo "test ok"
|
||||
|
|
@ -1,76 +0,0 @@
|
|||
package registry
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"ollie/pkg/tools"
|
||||
"ollie/pkg/tools/execute"
|
||||
)
|
||||
|
||||
// Server implements tools.Server backed by a session-aware registry.
|
||||
type Server struct {
|
||||
registry *Registry
|
||||
sessionID string
|
||||
execServer *execute.Server
|
||||
}
|
||||
|
||||
func NewServer(registry *Registry, sessionID string, execServer *execute.Server) *Server {
|
||||
return &Server{
|
||||
registry: registry,
|
||||
sessionID: sessionID,
|
||||
execServer: execServer,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) ListTools() ([]tools.ToolInfo, error) {
|
||||
return s.registry.Loaded(s.sessionID), nil
|
||||
}
|
||||
|
||||
func (s *Server) CallTool(ctx context.Context, tool string, args json.RawMessage) (json.RawMessage, error) {
|
||||
if _, promoted := s.registry.Lookup(s.sessionID, tool); !promoted {
|
||||
return nil, fmt.Errorf("tool_not_loaded: %s (must be loaded via /tools/load first)", tool)
|
||||
}
|
||||
|
||||
// Read the tool script
|
||||
script, err := execute.ReadTool(tool)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read tool %s: %w", tool, err)
|
||||
}
|
||||
|
||||
// Convert JSON object args to positional string args.
|
||||
// Try to extract an "args" array first (legacy call_tool format),
|
||||
// otherwise convert object values to positional strings.
|
||||
var positional []string
|
||||
var argMap map[string]interface{}
|
||||
if err := json.Unmarshal(args, &argMap); err == nil {
|
||||
if argsArr, ok := argMap["args"]; ok {
|
||||
if arr, ok := argsArr.([]interface{}); ok {
|
||||
for _, v := range arr {
|
||||
positional = append(positional, fmt.Sprintf("%v", v))
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Convert object values to positional strings (schema-defined order)
|
||||
for _, v := range argMap {
|
||||
positional = append(positional, fmt.Sprintf("%v", v))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Run the script through the execute kernel
|
||||
lang := execute.DetectLanguage(script)
|
||||
code := execute.InjectArgs(lang, tool, positional, script)
|
||||
result, err := s.execServer.Execute(ctx, code, lang, 30, "default", true)
|
||||
if err != nil {
|
||||
return json.Marshal(map[string]interface{}{
|
||||
"isError": true,
|
||||
"content": []map[string]string{{"type": "text", "text": err.Error()}},
|
||||
})
|
||||
}
|
||||
|
||||
return json.Marshal(map[string]interface{}{
|
||||
"content": []map[string]string{{"type": "text", "text": result}},
|
||||
})
|
||||
}
|
||||
Reference in New Issue