387 lines
11 KiB
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
|
|
}
|