core: support dynamic tool registration
This commit is contained in:
parent
3fb5528c97
commit
f08e0abdf0
|
|
@ -625,8 +625,6 @@ func TestWrapCommand_SortTiebreaker(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
func TestWrapCommand_ColonSeparatedPaths(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
dir1 := filepath.Join(tmpDir, "skills1")
|
||||
|
|
@ -692,8 +690,6 @@ func TestCheckPath_ColonSeparatedPaths(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
// ---- helpers ----
|
||||
|
||||
func mustWrapCommand(t *testing.T, cfg *Config, cmd []string, cwd string) []string {
|
||||
|
|
|
|||
|
|
@ -95,8 +95,6 @@ func WrapCommand(cfg *Config, originalCmd []string, cwd string, getenv EnvFunc)
|
|||
return args, nil
|
||||
}
|
||||
|
||||
|
||||
|
||||
// pathExists checks if a file or directory exists
|
||||
func pathExists(path string) bool {
|
||||
_, err := os.Stat(path)
|
||||
|
|
|
|||
|
|
@ -4,10 +4,10 @@ import "os"
|
|||
|
||||
// defaultCompactionModels maps backend names to cheap models suitable for compaction.
|
||||
var defaultCompactionModels = map[string]string{
|
||||
"anthropic": "claude-3-5-haiku-latest",
|
||||
"openai": "gpt-4o-mini",
|
||||
"openrouter": "deepseek-v4-flash",
|
||||
"gemini": "gemini-2.0-flash",
|
||||
"anthropic": "claude-3-5-haiku-latest",
|
||||
"openai": "gpt-4o-mini",
|
||||
"openrouter": "deepseek-v4-flash",
|
||||
"gemini": "gemini-2.0-flash",
|
||||
}
|
||||
|
||||
// backendNamer is the subset of backend.Backend needed for compaction model resolution.
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ type mockBackendForCompact struct {
|
|||
model string
|
||||
}
|
||||
|
||||
func (b *mockBackendForCompact) Name() string { return b.name }
|
||||
func (b *mockBackendForCompact) Name() string { return b.name }
|
||||
func (b *mockBackendForCompact) Model() string { return b.model }
|
||||
|
||||
func TestResolveCompactionModel_ConfigWins(t *testing.T) {
|
||||
|
|
|
|||
|
|
@ -104,8 +104,20 @@ func BuildRuntime(cfg *config.Config, d tools.Dispatcher, cwd string, env []stri
|
|||
}
|
||||
|
||||
exec := func(ctx context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error) {
|
||||
server, ok := serverOf[name]
|
||||
if !ok {
|
||||
// Dynamically resolve the server for this tool — handles
|
||||
// lazy-promoted tools that appeared after BuildRuntime.
|
||||
infos, listErr := d.ListTools()
|
||||
if listErr != nil {
|
||||
return "", nil, listErr
|
||||
}
|
||||
server := ""
|
||||
for _, t := range infos {
|
||||
if t.Name == name {
|
||||
server = t.Server
|
||||
break
|
||||
}
|
||||
}
|
||||
if server == "" {
|
||||
return "", nil, fmt.Errorf("unknown tool: %s", name)
|
||||
}
|
||||
raw, err := d.Dispatch(ctx, server, name, args)
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ func TestResultCache_HitSkipsExec(t *testing.T) {
|
|||
c.runtime.ClassifyTool = func(string) bool { return true }
|
||||
c.runtime.Exec = func(_ context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
||||
atomic.AddInt32(&execCount, 1)
|
||||
return "contents of a.txt", nil, nil
|
||||
return "contents of a.txt", nil, nil
|
||||
}
|
||||
|
||||
evs := collectEvents(context.Background(), c, "dup read")
|
||||
|
|
@ -56,7 +56,7 @@ func TestResultCache_DifferentArgsMiss(t *testing.T) {
|
|||
c.runtime.ClassifyTool = func(string) bool { return true }
|
||||
c.runtime.Exec = func(_ context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error) {
|
||||
atomic.AddInt32(&execCount, 1)
|
||||
return "result for " + string(args), nil, nil
|
||||
return "result for " + string(args), nil, nil
|
||||
}
|
||||
|
||||
collectEvents(context.Background(), c, "different args")
|
||||
|
|
@ -80,7 +80,7 @@ func TestResultCache_SerialToolNotCached(t *testing.T) {
|
|||
c.runtime.ClassifyTool = func(name string) bool { return false } // all serial
|
||||
c.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
||||
atomic.AddInt32(&execCount, 1)
|
||||
return "ok", nil, nil
|
||||
return "ok", nil, nil
|
||||
}
|
||||
|
||||
collectEvents(context.Background(), c, "serial tool")
|
||||
|
|
@ -101,10 +101,10 @@ func multiTurnToolsStream(turns [][]backend.ToolCall) func(context.Context, []ba
|
|||
ch := make(chan backend.StreamEvent, 1)
|
||||
ch <- backend.StreamEvent{ToolCalls: turns[i], Done: true, StopReason: "tool_calls"}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
return ch, nil
|
||||
|
||||
}
|
||||
return textStream("done"), nil
|
||||
return textStream("done"), nil
|
||||
|
||||
}
|
||||
}
|
||||
|
|
@ -128,7 +128,7 @@ func TestResultCache_ErrorNotCached(t *testing.T) {
|
|||
if n == 1 {
|
||||
return "", nil, &mockErr{"transient failure"}
|
||||
}
|
||||
return "ok now", nil, nil
|
||||
return "ok now", nil, nil
|
||||
}
|
||||
|
||||
evs := collectEvents(context.Background(), c, "error then ok")
|
||||
|
|
@ -160,7 +160,7 @@ func TestResultCache_IdenticalParallelReads(t *testing.T) {
|
|||
c.runtime.ClassifyTool = func(string) bool { return true } // both parallel-safe → same batch
|
||||
c.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
||||
atomic.AddInt32(&execCount, 1)
|
||||
return "file contents", nil, nil
|
||||
return "file contents", nil, nil
|
||||
}
|
||||
|
||||
evs := collectEvents(context.Background(), c, "identical parallel")
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ func alwaysFailStream() func(context.Context, []backend.Message, []backend.Tool,
|
|||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
||||
}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
return ch, nil
|
||||
|
||||
}
|
||||
}
|
||||
|
|
@ -63,7 +63,7 @@ func TestConsecutiveErrors_SoftLimitNudge(t *testing.T) {
|
|||
if n > consecutiveErrorSoftLimit {
|
||||
for _, m := range msgs {
|
||||
if m.Role == "user" && strings.Contains(m.Content, "your last several tool calls all failed") {
|
||||
return textStream("giving up"), nil
|
||||
return textStream("giving up"), nil
|
||||
|
||||
}
|
||||
}
|
||||
|
|
@ -76,7 +76,7 @@ return textStream("giving up"), nil
|
|||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
||||
}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
return ch, nil
|
||||
|
||||
}
|
||||
ch := make(chan backend.StreamEvent, 1)
|
||||
|
|
@ -87,7 +87,7 @@ return ch, nil
|
|||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
||||
}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
return ch, nil
|
||||
|
||||
}
|
||||
|
||||
|
|
@ -128,11 +128,11 @@ func TestConsecutiveErrors_ResetOnSuccess(t *testing.T) {
|
|||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
||||
}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
return ch, nil
|
||||
|
||||
}
|
||||
if n >= 11 {
|
||||
return textStream("done"), nil
|
||||
return textStream("done"), nil
|
||||
|
||||
}
|
||||
ch := make(chan backend.StreamEvent, 1)
|
||||
|
|
@ -143,14 +143,14 @@ return textStream("done"), nil
|
|||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
||||
}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
return ch, nil
|
||||
|
||||
}
|
||||
|
||||
c := newCore(t, be, nil)
|
||||
c.runtime.Exec = func(_ context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
||||
if name == "good_tool" {
|
||||
return "ok", nil, nil
|
||||
return "ok", nil, nil
|
||||
}
|
||||
return "", nil, fmt.Errorf("fails")
|
||||
}
|
||||
|
|
|
|||
|
|
@ -180,5 +180,5 @@ type taskStateState struct {
|
|||
ts *TaskState
|
||||
}
|
||||
|
||||
func (s *taskStateState) taskState() *TaskState { return s.ts }
|
||||
func (s *taskStateState) taskState() *TaskState { return s.ts }
|
||||
func (s *taskStateState) updateTaskState(ts TaskState) { s.ts = &ts }
|
||||
|
|
|
|||
|
|
@ -21,10 +21,10 @@ func toolsStream(calls []backend.ToolCall) func(context.Context, []backend.Messa
|
|||
ch := make(chan backend.StreamEvent, 1)
|
||||
ch <- backend.StreamEvent{ToolCalls: calls, Done: true, StopReason: "tool_calls"}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
return ch, nil
|
||||
|
||||
}
|
||||
return textStream("done"), nil
|
||||
return textStream("done"), nil
|
||||
|
||||
}
|
||||
}
|
||||
|
|
@ -57,7 +57,7 @@ func TestParallel_ConcurrentExecution(t *testing.T) {
|
|||
started <- struct{}{}
|
||||
select {
|
||||
case <-gate:
|
||||
return name + "-result", nil, nil
|
||||
return name + "-result", nil, nil
|
||||
case <-ctx.Done():
|
||||
return "", nil, ctx.Err()
|
||||
}
|
||||
|
|
@ -99,7 +99,7 @@ func TestParallel_SerialToolBreaksBatch(t *testing.T) {
|
|||
mu.Lock()
|
||||
order = append(order, name)
|
||||
mu.Unlock()
|
||||
return name + "-result", nil, nil
|
||||
return name + "-result", nil, nil
|
||||
}
|
||||
|
||||
collectEvents(context.Background(), c, "mixed tools")
|
||||
|
|
@ -181,7 +181,7 @@ func TestParallel_NilClassifyToolIsSerial(t *testing.T) {
|
|||
mu.Lock()
|
||||
order = append(order, name)
|
||||
mu.Unlock()
|
||||
return name + "-result", nil, nil
|
||||
return name + "-result", nil, nil
|
||||
}
|
||||
|
||||
evs := collectEvents(context.Background(), c, "serial fallback")
|
||||
|
|
|
|||
|
|
@ -18,11 +18,11 @@ func TestTextToolCall_ParsedAndExecuted(t *testing.T) {
|
|||
for _, m := range msgs {
|
||||
if m.Role == "tool" {
|
||||
// Second call after tool result: respond normally
|
||||
return textStream("done"), nil
|
||||
return textStream("done"), nil
|
||||
|
||||
}
|
||||
}
|
||||
return textStream(`call_tool: {"calls":[{"tool":"file_read","args":["/tmp/test.go"]}]}`), nil
|
||||
return textStream(`call_tool: {"calls":[{"tool":"file_read","args":["/tmp/test.go"]}]}`), nil
|
||||
|
||||
}
|
||||
|
||||
|
|
@ -37,7 +37,7 @@ return textStream(`call_tool: {"calls":[{"tool":"file_read","args":["/tmp/test.g
|
|||
if name != "call_tool" {
|
||||
t.Errorf("expected tool name 'call_tool', got %q", name)
|
||||
}
|
||||
return "file contents here", nil, nil
|
||||
return "file contents here", nil, nil
|
||||
}
|
||||
|
||||
evs := collectEvents(context.Background(), c, "read a file")
|
||||
|
|
@ -57,11 +57,11 @@ func TestTextToolCall_RelaxedJSON(t *testing.T) {
|
|||
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("done"), nil
|
||||
|
||||
}
|
||||
}
|
||||
return textStream(`execute_code: steps=[{code: "date"}]`), nil
|
||||
return textStream(`execute_code: steps=[{code: "date"}]`), nil
|
||||
|
||||
}
|
||||
|
||||
|
|
@ -76,7 +76,7 @@ return textStream(`execute_code: steps=[{code: "date"}]`), nil
|
|||
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
|
||||
return "Sat May 23 11:00:00 UTC 2026", nil, nil
|
||||
}
|
||||
|
||||
collectEvents(context.Background(), c, "what time is it")
|
||||
|
|
@ -91,7 +91,7 @@ return "Sat May 23 11:00:00 UTC 2026", nil, nil
|
|||
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
|
||||
return textStream("Here is how execute_code works in general."), nil
|
||||
|
||||
}
|
||||
|
||||
|
|
@ -115,7 +115,7 @@ return textStream("Here is how execute_code works in general."), nil
|
|||
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
|
||||
return textStream(`execute_code: steps=[{code: "print('hi')"}]`), nil
|
||||
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,6 @@ import (
|
|||
"ollie/pkg/tools"
|
||||
)
|
||||
|
||||
|
||||
// Runtime holds the swappable per-agent configuration. It contains everything
|
||||
// that changes on an /agent switch but is stable across turns within the same
|
||||
// agent. The agent struct stores a pointer to the active Runtime; switching
|
||||
|
|
@ -22,7 +21,7 @@ type Runtime struct {
|
|||
ClassifyTool toolClassifier
|
||||
ClassifyTier func(string, json.RawMessage) ResultTier
|
||||
GenParams backend.GenerationParams
|
||||
MaxSteps int
|
||||
MaxSteps int
|
||||
// CfgBackend and CfgModel are the backend/model overrides from the agent
|
||||
// config JSON. Empty means no override. Applied by the caller after
|
||||
// BuildRuntime returns.
|
||||
|
|
|
|||
|
|
@ -71,12 +71,12 @@ type anthropicCacheCtrl struct {
|
|||
}
|
||||
|
||||
type anthropicThinking struct {
|
||||
Type string `json:"type"` // always "enabled"
|
||||
Type string `json:"type"` // always "enabled"
|
||||
BudgetTokens int `json:"budget_tokens"`
|
||||
}
|
||||
|
||||
type anthropicMessage struct {
|
||||
Role string `json:"role"`
|
||||
Role string `json:"role"`
|
||||
Content []anthropicContentBlock `json:"content"`
|
||||
}
|
||||
|
||||
|
|
@ -314,8 +314,8 @@ func streamAnthropicSSE(body io.Reader, ch chan<- StreamEvent) {
|
|||
var v struct {
|
||||
Message struct {
|
||||
Usage struct {
|
||||
InputTokens int `json:"input_tokens"`
|
||||
CacheReadInputTokens int `json:"cache_read_input_tokens"`
|
||||
InputTokens int `json:"input_tokens"`
|
||||
CacheReadInputTokens int `json:"cache_read_input_tokens"`
|
||||
CacheCreationInputTokens int `json:"cache_creation_input_tokens"`
|
||||
} `json:"usage"`
|
||||
} `json:"message"`
|
||||
|
|
|
|||
|
|
@ -1051,10 +1051,10 @@ func kiroTokenExpiringSoon(expiresAt string) bool {
|
|||
// We probe OIDC first; if that key is absent we fall back to social.
|
||||
|
||||
const (
|
||||
kiroSQLiteProfileQuery = "SELECT value FROM state WHERE key = 'api.codewhisperer.profile' LIMIT 1;"
|
||||
kiroSQLiteOIDCTokenQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:odic:token' LIMIT 1;"
|
||||
kiroSQLiteDeviceRegQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:odic:device-registration' LIMIT 1;"
|
||||
kiroSQLiteSocialTokenQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:social:token' LIMIT 1;"
|
||||
kiroSQLiteProfileQuery = "SELECT value FROM state WHERE key = 'api.codewhisperer.profile' LIMIT 1;"
|
||||
kiroSQLiteOIDCTokenQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:odic:token' LIMIT 1;"
|
||||
kiroSQLiteDeviceRegQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:odic:device-registration' LIMIT 1;"
|
||||
kiroSQLiteSocialTokenQuery = "SELECT value FROM auth_kv WHERE key = 'kirocli:social:token' LIMIT 1;"
|
||||
)
|
||||
|
||||
type kiroSQLiteProfileState struct {
|
||||
|
|
|
|||
|
|
@ -49,9 +49,9 @@ type kiroRefreshTokenRequest struct {
|
|||
}
|
||||
|
||||
type kiroRefreshTokenResponse struct {
|
||||
AccessToken string `json:"accessToken"`
|
||||
AccessToken string `json:"accessToken"`
|
||||
RefreshToken string `json:"refreshToken,omitempty"`
|
||||
ExpiresIn *int `json:"expiresIn,omitempty"`
|
||||
ExpiresIn *int `json:"expiresIn,omitempty"`
|
||||
}
|
||||
|
||||
type kiroGenerateRequest struct {
|
||||
|
|
@ -90,11 +90,11 @@ type kiroChatMessage struct {
|
|||
}
|
||||
|
||||
type kiroUserInputMessage struct {
|
||||
Content string `json:"content"`
|
||||
UserInputMessageContext *kiroUserInputContext `json:"userInputMessageContext,omitempty"`
|
||||
Origin string `json:"origin,omitempty"`
|
||||
Images []kiroImageBlock `json:"images,omitempty"`
|
||||
ModelID string `json:"modelId,omitempty"`
|
||||
Content string `json:"content"`
|
||||
UserInputMessageContext *kiroUserInputContext `json:"userInputMessageContext,omitempty"`
|
||||
Origin string `json:"origin,omitempty"`
|
||||
Images []kiroImageBlock `json:"images,omitempty"`
|
||||
ModelID string `json:"modelId,omitempty"`
|
||||
}
|
||||
|
||||
type kiroImageBlock struct {
|
||||
|
|
@ -118,8 +118,8 @@ type kiroEnvState struct {
|
|||
}
|
||||
|
||||
type kiroAssistantResponseMessage struct {
|
||||
MessageID string `json:"messageId,omitempty"`
|
||||
Content string `json:"content"`
|
||||
MessageID string `json:"messageId,omitempty"`
|
||||
Content string `json:"content"`
|
||||
ToolUses []kiroToolUse `json:"toolUses,omitempty"`
|
||||
}
|
||||
|
||||
|
|
@ -144,14 +144,14 @@ type kiroToolUse struct {
|
|||
}
|
||||
|
||||
type kiroToolResult struct {
|
||||
ToolUseID string `json:"toolUseId"`
|
||||
Content []kiroToolResultContent `json:"content"`
|
||||
Status string `json:"status,omitempty"`
|
||||
ToolUseID string `json:"toolUseId"`
|
||||
Content []kiroToolResultContent `json:"content"`
|
||||
Status string `json:"status,omitempty"`
|
||||
}
|
||||
|
||||
type kiroToolResultContent struct {
|
||||
Text string `json:"text,omitempty"`
|
||||
JSON json.RawMessage `json:"json,omitempty"`
|
||||
Text string `json:"text,omitempty"`
|
||||
JSON json.RawMessage `json:"json,omitempty"`
|
||||
}
|
||||
|
||||
// stream event types
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ package backend
|
|||
import "os"
|
||||
|
||||
const (
|
||||
geminiBaseURL = "https://generativelanguage.googleapis.com/v1beta/openai/"
|
||||
geminiBaseURL = "https://generativelanguage.googleapis.com/v1beta/openai/"
|
||||
geminiDefaultModel = "gemini-2.5-flash-preview-05-20"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -107,10 +107,10 @@ func parseOllamaContextLength(r io.Reader) int {
|
|||
// -- wire types --
|
||||
|
||||
type ollamaMessage struct {
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
ToolCalls []ollamaToolCall `json:"tool_calls,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
ToolCalls []ollamaToolCall `json:"tool_calls,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||
}
|
||||
|
||||
type ollamaToolCall struct {
|
||||
|
|
|
|||
|
|
@ -47,9 +47,9 @@ func NewOpenAI(name, baseURL, apiKey string) (*OpenAIBackend, error) {
|
|||
return b, nil
|
||||
}
|
||||
|
||||
func (b *OpenAIBackend) Name() string { return b.name }
|
||||
func (b *OpenAIBackend) Model() string { return b.model }
|
||||
func (b *OpenAIBackend) SetModel(m string) { b.model = m; b.ctxLength = 0 }
|
||||
func (b *OpenAIBackend) Name() string { return b.name }
|
||||
func (b *OpenAIBackend) Model() string { return b.model }
|
||||
func (b *OpenAIBackend) SetModel(m string) { b.model = m; b.ctxLength = 0 }
|
||||
|
||||
func (b *OpenAIBackend) fetchModels(ctx context.Context) []openAIModelInfo {
|
||||
if len(b.cachedModels) > 0 {
|
||||
|
|
@ -208,10 +208,10 @@ type openAIResponseFmt struct {
|
|||
}
|
||||
|
||||
type openAIUsage struct {
|
||||
PromptTokens int `json:"prompt_tokens"`
|
||||
CompletionTokens int `json:"completion_tokens"`
|
||||
Cost float64 `json:"cost"` // OpenRouter only
|
||||
PromptTokensDetails struct {
|
||||
PromptTokens int `json:"prompt_tokens"`
|
||||
CompletionTokens int `json:"completion_tokens"`
|
||||
Cost float64 `json:"cost"` // OpenRouter only
|
||||
PromptTokensDetails struct {
|
||||
CachedTokens int `json:"cached_tokens"`
|
||||
} `json:"prompt_tokens_details"`
|
||||
}
|
||||
|
|
|
|||
|
|
@ -28,8 +28,8 @@ func (h *HookCmds) UnmarshalJSON(data []byte) error {
|
|||
// an array of strings (each element is a shell command whose stdout is
|
||||
// concatenated with newlines).
|
||||
type Prompt struct {
|
||||
Value []string // len==1 for a plain string; len>1 for command list
|
||||
IsExec bool // true when the value came from an array (execute mode)
|
||||
Value []string // len==1 for a plain string; len>1 for command list
|
||||
IsExec bool // true when the value came from an array (execute mode)
|
||||
}
|
||||
|
||||
func (p *Prompt) UnmarshalJSON(data []byte) error {
|
||||
|
|
@ -87,7 +87,7 @@ type Config struct {
|
|||
CompactionModel string `json:"compactionModel,omitempty"`
|
||||
// SystemPrompt overrides the embedded system prompt with a file path.
|
||||
// If set and the file exists, its content replaces the compiled-in default.
|
||||
SystemPrompt string `json:"systemPrompt,omitempty"`
|
||||
SystemPrompt string `json:"systemPrompt,omitempty"`
|
||||
}
|
||||
|
||||
// Load parses a Config from r.
|
||||
|
|
|
|||
|
|
@ -134,4 +134,4 @@ func TestToolsEnabled(t *testing.T) {
|
|||
if !cfg.ToolsEnabled() {
|
||||
t.Error("expected ToolsEnabled()=true")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -14,9 +14,9 @@ import (
|
|||
)
|
||||
|
||||
const (
|
||||
RequestTTL = 300 * time.Second
|
||||
FrameData = 'd'
|
||||
FrameExit = 'x'
|
||||
RequestTTL = 300 * time.Second
|
||||
FrameData = 'd'
|
||||
FrameExit = 'x'
|
||||
)
|
||||
|
||||
// NotifyFunc is called when a new request needs human attention.
|
||||
|
|
@ -25,9 +25,9 @@ type NotifyFunc func(req *Request)
|
|||
|
||||
// Broker manages elevation requests.
|
||||
type Broker struct {
|
||||
policy *PolicyStore
|
||||
notify NotifyFunc
|
||||
logf func(string, ...any)
|
||||
policy *PolicyStore
|
||||
notify NotifyFunc
|
||||
logf func(string, ...any)
|
||||
|
||||
mu sync.RWMutex
|
||||
pending map[string]*Request // id -> request
|
||||
|
|
@ -345,7 +345,7 @@ func (b *Broker) executeAndStream(conn net.Conn, cmd, cwd string, env map[string
|
|||
func (b *Broker) identifySession(conn net.Conn) string {
|
||||
// Use SO_PEERCRED to get caller PID, then map to session.
|
||||
// For now, return empty (can be wired up when session manager is available).
|
||||
uc, ok := conn.(*net.UnixConn);
|
||||
uc, ok := conn.(*net.UnixConn)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,16 +39,16 @@ func EnsureDefaults() {
|
|||
xdgRuntime = fmt.Sprintf("/run/user/%d", os.Getuid())
|
||||
}
|
||||
defaults := map[string]string{
|
||||
"OLLIE": filepath.Join(home, "mnt", "ollie"),
|
||||
"OLLIE_CFG_PATH": filepath.Join(home, ".config", "ollie"),
|
||||
"OLLIE_TOOLS_PATH": filepath.Join(home, ".config", "ollie", "tools"),
|
||||
"OLLIE_AGENTS_PATH": filepath.Join(home, ".config", "ollie", "agents"),
|
||||
"OLLIE_SKILLS_PATH": filepath.Join(home, ".config", "ollie", "skills"),
|
||||
"OLLIE_PROMPTS_PATH": filepath.Join(home, ".config", "ollie", "prompts"),
|
||||
"OLLIE_MEMORY_PATH": filepath.Join(home, ".config", "ollie", "memory"),
|
||||
"OLLIE_TMP_PATH": filepath.Join(home, ".local", "share", "ollie", "tmp"),
|
||||
"OLLIE_TRANSCRIPT_PATH": filepath.Join(home, ".config", "ollie", "transcript"),
|
||||
"OLLIE_ELEVATE_SOCKET": filepath.Join(xdgRuntime, "ollie", "elevate.sock"),
|
||||
"OLLIE": filepath.Join(home, "mnt", "ollie"),
|
||||
"OLLIE_CFG_PATH": filepath.Join(home, ".config", "ollie"),
|
||||
"OLLIE_TOOLS_PATH": filepath.Join(home, ".config", "ollie", "tools"),
|
||||
"OLLIE_AGENTS_PATH": filepath.Join(home, ".config", "ollie", "agents"),
|
||||
"OLLIE_SKILLS_PATH": filepath.Join(home, ".config", "ollie", "skills"),
|
||||
"OLLIE_PROMPTS_PATH": filepath.Join(home, ".config", "ollie", "prompts"),
|
||||
"OLLIE_MEMORY_PATH": filepath.Join(home, ".config", "ollie", "memory"),
|
||||
"OLLIE_TMP_PATH": filepath.Join(home, ".local", "share", "ollie", "tmp"),
|
||||
"OLLIE_TRANSCRIPT_PATH": filepath.Join(home, ".config", "ollie", "transcript"),
|
||||
"OLLIE_ELEVATE_SOCKET": filepath.Join(xdgRuntime, "ollie", "elevate.sock"),
|
||||
}
|
||||
for k, v := range defaults {
|
||||
if os.Getenv(k) == "" {
|
||||
|
|
|
|||
|
|
@ -33,9 +33,9 @@ var bootstrapTemplate string
|
|||
|
||||
// HostInfo holds environment details from the remote host.
|
||||
type HostInfo struct {
|
||||
Platform string `json:"platform"`
|
||||
Arch string `json:"arch"`
|
||||
IsGitRepo bool `json:"is_git_repo"`
|
||||
Platform string `json:"platform"`
|
||||
Arch string `json:"arch"`
|
||||
IsGitRepo bool `json:"is_git_repo"`
|
||||
}
|
||||
|
||||
// Server implements tools.Server by forwarding calls to a remote
|
||||
|
|
@ -392,9 +392,9 @@ func (s *Server) ping(ctx context.Context) error {
|
|||
// --- JSON-RPC types ---
|
||||
|
||||
type rpcRequest struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int64 `json:"id"`
|
||||
Method string `json:"method"`
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int64 `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -13,16 +13,16 @@ const ringBufSize = 64 * 1024 // 64KB ring buffer per detached process
|
|||
// DetachedProcess represents a process that the agent has detached from
|
||||
// but which continues running. The user can view its output and signal it.
|
||||
type DetachedProcess struct {
|
||||
PID int
|
||||
Command string
|
||||
Started time.Time
|
||||
Exited bool
|
||||
PID int
|
||||
Command string
|
||||
Started time.Time
|
||||
Exited bool
|
||||
ExitCode int
|
||||
|
||||
ring *ringBuffer
|
||||
cmd *os.Process
|
||||
done chan struct{}
|
||||
mu sync.Mutex
|
||||
ring *ringBuffer
|
||||
cmd *os.Process
|
||||
done chan struct{}
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// Info returns a plain-data snapshot of this process for external consumers.
|
||||
|
|
@ -114,4 +114,4 @@ func (r *ringBuffer) String() string {
|
|||
copy(out, r.buf[r.pos:])
|
||||
copy(out[r.size-r.pos:], r.buf[:r.pos])
|
||||
return string(out)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -69,11 +69,11 @@ type Server struct {
|
|||
blockedUntil time.Time
|
||||
|
||||
// Detached process management
|
||||
detachMu sync.Mutex
|
||||
detachCh chan struct{} // signal to detach the currently running process
|
||||
detached []*DetachedProcess
|
||||
OnDetach func(pid int, cmd string) // hook: called when a process is detached
|
||||
OnExit func(pid int, exitCode int) // hook: called when a detached process exits
|
||||
detachMu sync.Mutex
|
||||
detachCh chan struct{} // signal to detach the currently running process
|
||||
detached []*DetachedProcess
|
||||
OnDetach func(pid int, cmd string) // hook: called when a process is detached
|
||||
OnExit func(pid int, exitCode int) // hook: called when a detached process exits
|
||||
}
|
||||
|
||||
// Option configures a Server.
|
||||
|
|
@ -398,6 +398,7 @@ func (e *Server) SetEnv(key, value string) {
|
|||
e.OnEnvSet(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
// Close is called when the session ends. Calls OnClose hook if registered.
|
||||
func (e *Server) Close() {
|
||||
e.cleanupDetached()
|
||||
|
|
@ -406,7 +407,6 @@ func (e *Server) Close() {
|
|||
}
|
||||
}
|
||||
|
||||
|
||||
// executeElevated runs cmd outside the sandbox via the integrated elevation broker.
|
||||
// Connects to the broker socket, sends the request with the current env,
|
||||
// and streams the framed response back.
|
||||
|
|
|
|||
|
|
@ -8,19 +8,19 @@ import (
|
|||
// 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",
|
||||
"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,
|
||||
|
|
|
|||
|
|
@ -76,6 +76,11 @@ func detectLanguage(code string) string {
|
|||
return "bash"
|
||||
}
|
||||
|
||||
// DetectLanguage is the exported version of detectLanguage.
|
||||
func DetectLanguage(code string) string {
|
||||
return detectLanguage(code)
|
||||
}
|
||||
|
||||
// injectArgs prepends language-appropriate argument binding to code.
|
||||
func injectArgs(language, name string, args []string, code string) string {
|
||||
switch language {
|
||||
|
|
@ -141,6 +146,11 @@ func injectArgs(language, name string, args []string, code string) string {
|
|||
}
|
||||
}
|
||||
|
||||
// InjectArgs is the exported version of injectArgs.
|
||||
func InjectArgs(language, name string, args []string, code string) string {
|
||||
return injectArgs(language, name, args, code)
|
||||
}
|
||||
|
||||
// ansiCEscape escapes a string for embedding in a bash $'...' literal.
|
||||
func ansiCEscape(s string) string {
|
||||
var b strings.Builder
|
||||
|
|
@ -250,6 +260,11 @@ func extractShortDescription(prompt string) string {
|
|||
return ""
|
||||
}
|
||||
|
||||
// ExtractShortDescription is the exported version of extractShortDescription.
|
||||
func ExtractShortDescription(prompt string) string {
|
||||
return extractShortDescription(prompt)
|
||||
}
|
||||
|
||||
// ToolPrompt returns the full ollie:prompt content for a named tool.
|
||||
// Returns "" if the tool has no prompt or doesn't exist.
|
||||
func ToolPrompt(name string) string {
|
||||
|
|
@ -280,4 +295,3 @@ func ReadTool(name string) (string, error) {
|
|||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,154 @@
|
|||
package registry
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"ollie/pkg/tools"
|
||||
"ollie/pkg/tools/execute"
|
||||
)
|
||||
|
||||
type Registry struct {
|
||||
mu sync.RWMutex
|
||||
global map[string]tools.ToolInfo
|
||||
sessions map[string]map[string]tools.ToolInfo
|
||||
revisions map[string]uint64
|
||||
}
|
||||
|
||||
func NewRegistry() (*Registry, error) {
|
||||
r := &Registry{
|
||||
global: make(map[string]tools.ToolInfo),
|
||||
sessions: make(map[string]map[string]tools.ToolInfo),
|
||||
revisions: make(map[string]uint64),
|
||||
}
|
||||
if err := r.Discover(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
|
||||
func (r *Registry) Discover() error {
|
||||
dir := execute.ToolsPath()
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read tools dir %s: %w", dir, err)
|
||||
}
|
||||
global := make(map[string]tools.ToolInfo)
|
||||
for _, e := range entries {
|
||||
if e.IsDir() || e.Name() == "idx" || e.Name()[0] == '.' {
|
||||
continue
|
||||
}
|
||||
data, err := os.ReadFile(filepath.Join(dir, e.Name()))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
script := string(data)
|
||||
global[e.Name()] = ParseToolInfo(e.Name(), script)
|
||||
}
|
||||
r.mu.Lock()
|
||||
r.global = global
|
||||
r.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Registry) Summaries() []tools.ToolInfo {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
var summaries []tools.ToolInfo
|
||||
for _, info := range r.global {
|
||||
summaries = append(summaries, tools.ToolInfo{
|
||||
Name: info.Name,
|
||||
Description: info.Description,
|
||||
})
|
||||
}
|
||||
return summaries
|
||||
}
|
||||
|
||||
func (r *Registry) Load(sessionID, name string) error {
|
||||
r.mu.RLock()
|
||||
tool, exists := r.global[name]
|
||||
r.mu.RUnlock()
|
||||
if !exists {
|
||||
return fmt.Errorf("tool not found: %s", name)
|
||||
}
|
||||
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if r.sessions[sessionID] == nil {
|
||||
r.sessions[sessionID] = make(map[string]tools.ToolInfo)
|
||||
}
|
||||
|
||||
if _, already := r.sessions[sessionID][name]; already {
|
||||
return nil
|
||||
}
|
||||
|
||||
r.sessions[sessionID][name] = tool
|
||||
r.revisions[sessionID]++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Registry) Unload(sessionID, name string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
sessionTools, ok := r.sessions[sessionID]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
if _, exists := sessionTools[name]; !exists {
|
||||
return nil
|
||||
}
|
||||
|
||||
delete(sessionTools, name)
|
||||
r.revisions[sessionID]++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Registry) Loaded(sessionID string) []tools.ToolInfo {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
sessionTools, ok := r.sessions[sessionID]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
var loaded []tools.ToolInfo
|
||||
for _, info := range sessionTools {
|
||||
loaded = append(loaded, info)
|
||||
}
|
||||
return loaded
|
||||
}
|
||||
|
||||
func (r *Registry) Lookup(sessionID, name string) (tools.ToolInfo, bool) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
sessionTools, ok := r.sessions[sessionID]
|
||||
if !ok {
|
||||
return tools.ToolInfo{}, false
|
||||
}
|
||||
|
||||
tool, exists := sessionTools[name]
|
||||
return tool, exists
|
||||
}
|
||||
|
||||
func (r *Registry) Revision(sessionID string) uint64 {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.revisions[sessionID]
|
||||
}
|
||||
|
||||
func (r *Registry) GlobalToolNames() []string {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
names := make([]string, 0, len(r.global))
|
||||
for name := range r.global {
|
||||
names = append(names, name)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
|
@ -0,0 +1,159 @@
|
|||
package registry
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
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{
|
||||
registry: r,
|
||||
sessionID: 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 tools[0].Name != names[0] {
|
||||
t.Fatalf("expected %s, got %s", names[0], tools[0].Name)
|
||||
}
|
||||
if tools[0].InputSchema == nil {
|
||||
t.Fatal("expected InputSchema to be non-nil")
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,71 @@
|
|||
package registry
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"ollie/pkg/tools"
|
||||
"ollie/pkg/tools/execute"
|
||||
)
|
||||
|
||||
type ToolMeta struct {
|
||||
ReadOnly bool `json:"readOnly,omitempty"`
|
||||
RequiresApproval bool `json:"requiresApproval,omitempty"`
|
||||
NetworkAccess bool `json:"networkAccess,omitempty"`
|
||||
PathScope string `json:"pathScope,omitempty"`
|
||||
SandboxLevel string `json:"sandboxLevel,omitempty"`
|
||||
Examples []string `json:"examples,omitempty"`
|
||||
Tags []string `json:"tags,omitempty"`
|
||||
}
|
||||
|
||||
func ExtractArgsSchema(script string) json.RawMessage {
|
||||
for _, line := range strings.Split(script, "\n") {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if strings.HasPrefix(trimmed, "# args_json:") {
|
||||
schema := strings.TrimSpace(strings.TrimPrefix(trimmed, "# args_json:"))
|
||||
if schema != "" {
|
||||
return json.RawMessage(schema)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ExtractReturnSchema(script string) json.RawMessage {
|
||||
for _, line := range strings.Split(script, "\n") {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if strings.HasPrefix(trimmed, "# returns_json:") {
|
||||
schema := strings.TrimSpace(strings.TrimPrefix(trimmed, "# returns_json:"))
|
||||
if schema != "" {
|
||||
return json.RawMessage(schema)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ExtractMetadata(script string) ToolMeta {
|
||||
var meta ToolMeta
|
||||
for _, line := range strings.Split(script, "\n") {
|
||||
trimmed := strings.TrimSpace(line)
|
||||
if strings.Contains(trimmed, "ollie:parallel read") {
|
||||
meta.ReadOnly = true
|
||||
}
|
||||
}
|
||||
return meta
|
||||
}
|
||||
|
||||
func ParseToolInfo(name, script string) tools.ToolInfo {
|
||||
prompt := execute.ExtractPrompt(script)
|
||||
desc := execute.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"]}`)
|
||||
}
|
||||
return tools.ToolInfo{
|
||||
Name: name,
|
||||
Description: desc,
|
||||
InputSchema: argsSchema,
|
||||
Prompt: prompt,
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,76 @@
|
|||
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}},
|
||||
})
|
||||
}
|
||||
|
|
@ -16,7 +16,7 @@ type ToolInfo struct {
|
|||
InputSchema json.RawMessage
|
||||
// Prompt is the usage documentation for this tool, extracted from
|
||||
// the script's ollie:prompt block. Included in the system prompt.
|
||||
Prompt string
|
||||
Prompt string
|
||||
}
|
||||
|
||||
// Server is the interface satisfied by any tool server.
|
||||
|
|
|
|||
Reference in New Issue