ollie/cmd/olliesrv/internal/agent/dispatch.go

387 lines
11 KiB
Go

// dispatch.go — Tool execution, batching, and conflict detection.
//
// execToolCalls processes a list of tool calls from a single LLM turn.
// Non-conflicting calls run in parallel (reads, writes to different paths).
// Writes to the same path and global-scope tools serialize.
package agent
import (
"context"
"encoding/json"
"errors"
"fmt"
"strings"
"sync"
"ollie/cmd/olliesrv/internal/backend"
toolclient "ollie/cmd/olliesrv/internal/toolclient"
)
// conflictKeyGlobal is a sentinel indicating a tool conflicts with everything.
const conflictKeyGlobal = "\x00GLOBAL"
// execToolCalls executes a list of tool calls with batching and conflict detection.
func (ag *Agent) execToolCalls(ctx context.Context, toolCalls []backend.ToolCall) ([]toolResult, bool) {
rt := ag.runtime
results := make([]toolResult, 0, len(toolCalls))
cancelledResult := func(call backend.ToolCall) toolResult {
return toolResult{
ToolCallID: call.ID,
Name: call.Name,
Content: `{"status":"cancelled","error":"interrupted"}`,
IsError: true,
}
}
fillCancelled := func(calls []backend.ToolCall) {
for _, c := range calls {
cr := cancelledResult(c)
ag.emit(Event{Role: "tool", Name: c.Name, Content: cr.Content, OutputFormat: toolOutputFormat(rt, c.Name)})
results = append(results, cr)
}
}
for i := 0; i < len(toolCalls); {
if ctx.Err() != nil {
fillCancelled(toolCalls[i:])
return results, true
}
// Build a batch of non-conflicting calls using greedy grouping.
batch := []backend.ToolCall{toolCalls[i]}
batchWrites := toolConflictKeys(rt, toolCalls[i])
j := i + 1
for j < len(toolCalls) {
// A global barrier seals the batch — nothing else can join.
if hasGlobal(batchWrites) {
break
}
keys := toolConflictKeys(rt, toolCalls[j])
if conflictsWithSet(keys, batchWrites) {
break
}
batch = append(batch, toolCalls[j])
batchWrites = mergeKeys(batchWrites, keys)
j++
}
if len(batch) == 1 {
tr, wasInt := ag.execOne(ctx, batch[0])
results = append(results, tr)
if wasInt {
fillCancelled(toolCalls[j:])
return results, true
}
} else {
batchResults, wasInt := ag.execBatch(ctx, batch)
results = append(results, batchResults...)
if wasInt {
fillCancelled(toolCalls[j:])
return results, true
}
}
i = j
}
return results, false
}
// toolConflictKeys returns the set of resource keys a tool call touches.
// Scope "read" tools return nil (never conflict). Scope "write" tools
// return their file path (or global if no path detectable). Everything
// else (scope "global" or unset) is a full serialization barrier.
func toolConflictKeys(rt *Runtime, call backend.ToolCall) []string {
scope := rt.ToolMeta[call.Name].Scope
if scope == "read" {
return nil // reads never conflict
}
if scope == "write" {
if p := extractFilePath(call.Arguments); p != "" {
return []string{p}
}
return []string{conflictKeyGlobal}
}
// "global", empty, or anything else: full barrier
return []string{conflictKeyGlobal}
}
// conflictsWithSet returns true if keys conflicts with the accumulated set.
// nil keys (scope "read") never conflict. Global sentinel conflicts with any
// non-empty set. Path keys conflict if the same path exists in the set.
func conflictsWithSet(keys []string, set []string) bool {
if len(keys) == 0 {
return false // scope "read": never conflicts
}
if len(set) == 0 {
return false // nothing accumulated yet
}
for _, k := range keys {
if k == conflictKeyGlobal {
return true // global conflicts with any non-empty set
}
for _, s := range set {
if s == conflictKeyGlobal || s == k {
return true
}
}
}
return false
}
// mergeKeys appends keys into set (no dedup needed for small sets).
func mergeKeys(set, keys []string) []string {
return append(set, keys...)
}
// hasGlobal returns true if set contains the global barrier sentinel.
func hasGlobal(set []string) bool {
for _, s := range set {
if s == conflictKeyGlobal {
return true
}
}
return false
}
// execBatch fans out a batch of parallel-safe calls, deduplicating identical ones.
func (ag *Agent) execBatch(ctx context.Context, batch []backend.ToolCall) ([]toolResult, bool) {
rt := ag.runtime
type inflightResult struct {
tr toolResult
wasInt bool
}
inflight := make(map[string]int) // key → index of first occurrence
batchResults := make([]toolResult, len(batch))
uniqueResults := make([]inflightResult, len(batch))
var wg sync.WaitGroup
sem := make(chan struct{}, maxParallelToolCalls)
for k, call := range batch {
key := call.Name + "\x00" + string(call.Arguments)
if _, dup := inflight[key]; dup {
continue
}
inflight[key] = k
wg.Add(1)
go func(k int, call backend.ToolCall) {
defer wg.Done()
select {
case sem <- struct{}{}:
case <-ctx.Done():
uniqueResults[k].wasInt = true
return
}
defer func() { <-sem }()
uniqueResults[k].tr, uniqueResults[k].wasInt = ag.execOne(ctx, call)
}(k, call)
}
wg.Wait()
var interrupted bool
for k, call := range batch {
key := call.Name + "\x00" + string(call.Arguments)
first := inflight[key]
if k == first {
batchResults[k] = uniqueResults[k].tr
if uniqueResults[k].wasInt {
interrupted = true
}
} else {
batchResults[k] = toolResult{
ToolCallID: call.ID,
Name: call.Name,
Content: uniqueResults[first].tr.Content,
IsError: uniqueResults[first].tr.IsError,
}
ag.emit(Event{Role: "call", Name: call.Name, Content: string(call.Arguments)})
ag.emit(Event{Role: "tool", Name: call.Name, Content: batchResults[k].Content, OutputFormat: toolOutputFormat(rt, call.Name)})
}
}
return batchResults, interrupted
}
// execOne runs a single tool call and reports whether the context was cancelled.
func (ag *Agent) execOne(ctx context.Context, call backend.ToolCall) (toolResult, bool) {
rt := ag.runtime
resultCache := &ag.resultCache
if call.Name == "" {
return toolResult{ToolCallID: call.ID, Name: call.Name, Content: "error: empty tool name", IsError: true}, false
}
if ctx.Err() != nil {
cr := toolResult{ToolCallID: call.ID, Name: call.Name, Content: `{"status":"cancelled","error":"interrupted"}`, IsError: true}
ag.emit(Event{Role: "tool", Name: call.Name, Content: cr.Content, OutputFormat: toolOutputFormat(rt, call.Name)})
return cr, true
}
ag.emit(Event{Role: "call", Name: call.Name, Content: string(call.Arguments)})
readSafe := rt.ToolMeta[call.Name].Scope == "read"
if readSafe {
key := call.Name + "\x00" + string(call.Arguments)
if entry, ok := resultCache.Load(key); ok {
if cacheValid(entry, call.Arguments) {
ag.emit(Event{Role: "tool", Name: call.Name, Content: entry.Result, OutputFormat: toolOutputFormat(rt, call.Name)})
return toolResult{ToolCallID: call.ID, Name: call.Name, Content: entry.Result}, false
}
resultCache.Delete(key)
}
}
var result string
var resultBlocks []backend.ContentBlock
var isErr bool
toolServer := rt.ToolServer
// Check for background flag — if present, execute via proc/new.bg
if isBackground(call.Arguments) {
strippedArgs := stripBackgroundFlag(call.Arguments)
if toolServer == nil {
result = "error: no tool server available for background execution"
isErr = true
} else {
// Extract command description for display
cmdDesc := extractCmdDesc(call.Name, call.Arguments)
pid, err := toolServer.CallToolBackground(call.Name, strippedArgs)
if err != nil {
result = fmt.Sprintf("error: background exec: %v", err)
isErr = true
} else {
result = fmt.Sprintf("<system-proc-background>\nid=%d cmd=%q\n</system-proc-background>", pid, cmdDesc)
// Notify proc.start via callback
if ag.onProcStart != nil {
ag.onProcStart(ag.id, pid, call.Name, cmdDesc)
}
}
}
ag.emit(Event{Role: "tool", Name: call.Name, Content: result, OutputFormat: toolOutputFormat(rt, call.Name)})
return toolResult{ToolCallID: call.ID, Name: call.Name, Content: result, IsError: isErr}, false
}
if rt.Exec != nil {
out, blocks, err := rt.Exec(ctx, call.Name, call.Arguments)
if err != nil {
isErr = true
if ctx.Err() != nil {
result = "error: tool execution interrupted by user"
if injected := ag.popInject(); injected != "" {
result += "\n\n<system-user-interruption>\n" + injected + "\n</system-user-interruption>"
}
ag.emit(Event{Role: "tool", Name: call.Name, Content: result, OutputFormat: toolOutputFormat(rt, call.Name)})
return toolResult{ToolCallID: call.ID, Name: call.Name, Content: result, IsError: true}, true
}
var rlErr *toolclient.RateLimitedError
if errors.As(err, &rlErr) {
result = fmt.Sprintf("error: shell is rate-limited — blocked for %v. Do not call shell() again until the block expires. Use other tools or wait.", rlErr.Remaining)
} else {
result = fmt.Sprintf("error: %v", err)
}
} else {
result = out
resultBlocks = blocks
}
} else {
result = "error: no tool executor configured"
isErr = true
}
// Append suffix (user-interruptions, truncation).
var suffix string
if injected := ag.popInject(); injected != "" {
suffix += "\n\n<system-user-interruption>\n" + injected + "\n</system-user-interruption>"
}
result += suffix
if len(result) > defaultToolResultMaxBytes {
orig := len(result)
result = strings.ToValidUTF8(result[:defaultToolResultMaxBytes], "")
suffix = fmt.Sprintf("\n\n[HARD LIMIT: %s output truncated — %d of %d bytes shown. This is a safety ceiling, not a semantic boundary.]",
call.Name, defaultToolResultMaxBytes, orig)
result += suffix
}
if readSafe && !isErr {
path := extractFilePath(call.Arguments)
mtime, size := fileStat(path)
resultCache.Store(call.Name+"\x00"+string(call.Arguments), cachedResult{
Result: result,
ModTime: mtime,
Size: size,
})
}
ag.emit(Event{Role: "tool", Name: call.Name, Content: result, OutputFormat: toolOutputFormat(rt, call.Name)})
tier := MemoryHot
if !isErr {
tier = toolMemoryTier(rt, call.Name)
}
return toolResult{ToolCallID: call.ID, Name: call.Name, Content: result, ContentBlocks: resultBlocks, IsError: isErr, Tier: tier}, false
}
// popInject returns and clears any pending inject text.
func (ag *Agent) popInject() string {
if p := ag.pendingInject.Swap(nil); p != nil {
return *p
}
return ""
}
// --- Background execution helpers ---
// isBackground returns true if the tool args contain "background": true.
func isBackground(args json.RawMessage) bool {
var m map[string]json.RawMessage
if json.Unmarshal(args, &m) != nil {
return false
}
raw, ok := m["background"]
if !ok {
return false
}
var v bool
if json.Unmarshal(raw, &v) == nil {
return v
}
// Also accept string "true"
var s string
if json.Unmarshal(raw, &s) == nil {
return s == "true" || s == "1"
}
return false
}
// stripBackgroundFlag removes the "background" key from JSON args.
func stripBackgroundFlag(args json.RawMessage) json.RawMessage {
var m map[string]json.RawMessage
if json.Unmarshal(args, &m) != nil {
return args
}
delete(m, "background")
out, _ := json.Marshal(m)
return out
}
// extractCmdDesc returns a short human-readable description of the tool call.
func extractCmdDesc(name string, args json.RawMessage) string {
var m map[string]json.RawMessage
if json.Unmarshal(args, &m) != nil {
return name
}
// For shell, use the cmd field
if raw, ok := m["cmd"]; ok {
var cmd string
if json.Unmarshal(raw, &cmd) == nil {
if len(cmd) > 80 {
cmd = cmd[:80] + "..."
}
return cmd
}
}
return name
}