ollie/agent/loop.go

706 lines
21 KiB
Go

package agent
import (
"context"
"encoding/json"
"errors"
"fmt"
"os"
"strings"
"sync"
"time"
"ollie/backend"
"ollie/paths"
"ollie/toolsrv"
)
const maxTransientRetries = 3
// retryBaseDelay is the base delay for rate-limit retries (5<<attempt seconds).
// Tests can override this to speed up retry tests.
var retryBaseDelay = 5 * time.Second
// streamDropBaseDelay is the base delay for stream-drop retries (2<<attempt seconds).
var streamDropBaseDelay = 2 * time.Second
// Consecutive tool-error thresholds. At the soft limit the model is nudged
// to try a different approach; at the hard limit the loop aborts.
const (
consecutiveErrorSoftLimit = 5
consecutiveErrorHardLimit = 10
)
type toolExecutor func(ctx context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error)
// cachedResult stores a tool result with file metadata for staleness detection.
type cachedResult struct {
Result string
ModTime int64 // unix nanoseconds of the file at cache time; 0 if no file
Size int64 // file size at cache time; -1 if no file
}
// extractFilePath attempts to pull a "path" field from JSON tool args.
func extractFilePath(args json.RawMessage) string {
var m map[string]json.RawMessage
if json.Unmarshal(args, &m) != nil {
return ""
}
raw, ok := m["path"]
if !ok {
return ""
}
var p string
if json.Unmarshal(raw, &p) != nil {
return ""
}
return p
}
// cacheValid checks if a cached result is still valid by stat-ing the file.
func cacheValid(entry cachedResult, args json.RawMessage) bool {
if entry.Size < 0 {
return true
}
path := extractFilePath(args)
if path == "" {
return true
}
info, err := os.Stat(path)
if err != nil {
return false
}
return info.ModTime().UnixNano() == entry.ModTime && info.Size() == entry.Size
}
// fileStat returns (modtime nanoseconds, size) for a path, or (0, -1) if unavailable.
func fileStat(path string) (int64, int64) {
if path == "" {
return 0, -1
}
info, err := os.Stat(path)
if err != nil {
return 0, -1
}
return info.ModTime().UnixNano(), info.Size()
}
func toolOutputFormat(rt *Runtime, name string) string {
if ti, ok := rt.ToolMeta[name]; ok {
return ti.OutputFormat
}
return ""
}
func toolMemoryTier(rt *Runtime, name string) MemoryTier {
if ti, ok := rt.ToolMeta[name]; ok {
switch ti.Tier {
case "cold":
return MemoryCold
case "warm":
return MemoryWarm
}
}
return MemoryHot
}
// TurnCtx holds per-turn closures and state that vary between turns within
// the same agent. Everything else comes from *Runtime. Embeds context.Context
// so the turn's cancellation is carried with the turn state rather than as a
// separate parameter.
type TurnCtx struct {
context.Context
Output EventHandler
PopInject func() string
AutoCompact func(ctx context.Context)
Save func()
IncrToolCallCount func() int64
ResultCache *sync.Map
}
// errorState tracks consecutive and repeated tool errors across turns.
type errorState struct {
consecutive int
lastSig string
repeatCount int
}
func run(rt *Runtime, ctx TurnCtx, h *History) error {
var step int
var errs errorState
resultCache := ctx.ResultCache
if resultCache == nil {
resultCache = &sync.Map{}
}
for {
emit(ctx, Event{Role: "state", Content: "thinking"})
if budget := contextBudget(ctx, rt.Backend); budget > 0 && h.estimateTokens() > budget {
h.stripCold(ctx, rt.Backend)
}
history := h.history()
if preamble := rt.PreambleString(); preamble != "" {
history = append([]backend.Message{{Role: "system", Content: preamble}}, history...)
}
responseID := paths.NewUUID()
content, reasoning, toolCalls, stopReason, err := streamResponse(rt, ctx, history, responseID, step)
if err != nil {
// On context cancellation mid-stream, persist partial state.
if ctx.Err() != nil && (content != "" || len(toolCalls) > 0) {
msg := backend.Message{ID: responseID, Role: "assistant", Content: content, Reasoning: reasoning, ToolCalls: toolCalls}
var partialResults []toolResult
for _, tc := range toolCalls {
partialResults = append(partialResults, toolResult{
ToolCallID: tc.ID, Name: tc.Name,
Content: `{"status":"cancelled","error":"stream interrupted"}`, IsError: true,
})
}
h.update(msg, partialResults)
ctx.Save()
}
return err
}
switch stopReason {
case "stop", "tool_calls", "length", "error", "":
default:
return fmt.Errorf("step %d: %s", step, stopReason)
}
// Parse text-based tool calls for models without function calling.
if len(toolCalls) == 0 && len(rt.Tools) > 0 {
if parsed := parseTextToolCalls(content, rt.Tools); len(parsed) > 0 {
toolCalls = parsed
}
}
msg := backend.Message{ID: responseID, Role: "assistant", Content: content, Reasoning: reasoning, ToolCalls: toolCalls}
results, interrupted := execToolCalls(rt, ctx, toolCalls, resultCache)
h.update(msg, results)
ctx.Save()
if interrupted {
return ctx.Err()
}
if len(toolCalls) == 0 {
break
}
if abort := trackErrors(h, ctx, results, &errs, step); abort != nil {
return abort
}
// Reset step counter when an action-oriented tool succeeds.
if rt.ToolMeta != nil {
for _, r := range results {
if !r.IsError && rt.ToolMeta[r.Name].ResetsCounter {
step = 0
break
}
}
}
// Step budget warnings.
if rt.MaxSteps > 0 {
halfBudget := rt.MaxSteps / 2
if step == halfBudget {
h.update(backend.Message{
Role: "user",
Content: fmt.Sprintf("<system-step-budget-warning>\nYou have used %d/%d research steps without taking action. Consider making progress — write code, edit files, or run commands. Action tools reset this counter.\n</system-step-budget-warning>", step, rt.MaxSteps),
}, nil)
} else if step >= rt.MaxSteps-1 {
emit(ctx, Event{Role: "maxsteps", Content: fmt.Sprintf("%d", step+1)})
h.update(backend.Message{
Role: "user",
Content: fmt.Sprintf("<system-step-budget-stop>\nStep budget exhausted (%d/%d steps used). Stop calling tools. Summarize what you have done and what remains, then stop.\n</system-step-budget-stop>", step+1, rt.MaxSteps),
}, nil)
break
}
}
if ctx.AutoCompact != nil {
ctx.AutoCompact(ctx)
}
step++
}
return nil
}
// ── streamResponse ──────────────────────────────────────────────────────────
// streamResponse calls ChatStream with retry on rate limits, transient errors,
// and mid-stream drops. Returns the assembled content, reasoning, tool calls,
// and stop reason.
func streamResponse(rt *Runtime, ctx TurnCtx, history []backend.Message, responseID string, step int) (content string, reasoning string, toolCalls []backend.ToolCall, stopReason string, err error) {
var contentBuf strings.Builder
var reasoningBuf strings.Builder
for attempt := range maxTransientRetries + 1 {
contentBuf.Reset()
reasoningBuf.Reset()
toolCalls = nil
ch, streamErr := rt.Backend.ChatStream(ctx, history, rt.Tools, rt.GenParams)
if streamErr != nil {
if ctx.Err() != nil {
return "", "", nil, "", ctx.Err()
}
wait, retryable := transientWait(streamErr, attempt)
if !retryable || attempt >= maxTransientRetries {
return "", "", nil, "", fmt.Errorf("step %d: %w", step, streamErr)
}
var rlErr *backend.RateLimitError
if errors.As(streamErr, &rlErr) {
emit(ctx, Event{Role: "limitretry"})
}
if err := retryCountdown(ctx, wait); err != nil {
return "", "", nil, "", fmt.Errorf("step %d: %w", step, err)
}
continue
}
var done bool
var hadReasoning bool
for ev := range ch {
if ev.Reasoning != "" {
if !hadReasoning {
emit(ctx, Event{Role: "reasoning", Content: "<think>\n"})
hadReasoning = true
}
reasoningBuf.WriteString(ev.Reasoning)
emit(ctx, Event{Role: "reasoning", Content: ev.Reasoning})
}
if ev.Content != "" {
if hadReasoning {
emit(ctx, Event{Role: "reasoning", Content: "\n</think>\n"})
hadReasoning = false
}
contentBuf.WriteString(ev.Content)
emit(ctx, Event{Role: "assistant", Content: ev.Content, ResponseID: responseID})
}
toolCalls = append(toolCalls, ev.ToolCalls...)
if ev.Done {
if hadReasoning {
emit(ctx, Event{Role: "reasoning", Content: "\n</think>\n"})
hadReasoning = false
}
stopReason = ev.StopReason
done = true
if ev.Usage.InputTokens > 0 || ev.Usage.OutputTokens > 0 {
emit(ctx, Event{
Role: "usage",
Content: fmt.Sprintf("%d %d 0 %g %d %d", ev.Usage.InputTokens, ev.Usage.OutputTokens, ev.Usage.CostUSD, ev.Usage.CachedInputTokens, ev.Usage.CacheCreationTokens),
})
} else {
inChars := 0
for _, m := range history {
inChars += len(m.Content)
for _, tc := range m.ToolCalls {
inChars += len(tc.Name) + len(tc.Arguments)
}
}
emit(ctx, Event{
Role: "usage",
Content: fmt.Sprintf("%d %d 1 0", inChars/4, contentBuf.Len()/4),
})
}
break
}
}
if done {
return contentBuf.String(), reasoningBuf.String(), toolCalls, stopReason, nil
}
if ctx.Err() != nil {
if hadReasoning {
emit(ctx, Event{Role: "reasoning", Content: "\n</think>\n"})
}
// Return partial state for caller to persist.
return contentBuf.String(), reasoningBuf.String(), toolCalls, "", ctx.Err()
}
// Pure stream drop — retry.
if attempt >= maxTransientRetries {
return "", "", nil, "", fmt.Errorf("step %d: stream dropped (no more retries)", step)
}
if hadReasoning {
emit(ctx, Event{Role: "reasoning", Content: "\n</think>\n"})
}
wait := streamDropBaseDelay << attempt
if err := retryCountdown(ctx, wait); err != nil {
return "", "", nil, "", fmt.Errorf("step %d: %w", step, err)
}
}
return contentBuf.String(), reasoningBuf.String(), toolCalls, stopReason, nil
}
// ── execToolCalls ───────────────────────────────────────────────────────────
// execToolCalls dispatches all tool calls, running consecutive read-safe
// calls in parallel. Returns the results and whether the loop was interrupted.
func execToolCalls(rt *Runtime, ctx TurnCtx, toolCalls []backend.ToolCall, resultCache *sync.Map) ([]toolResult, bool) {
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)
emit(ctx, Event{Role: "tool", Name: c.Name, Content: cr.Content, OutputFormat: toolOutputFormat(rt, c.Name)})
results = append(results, cr)
}
}
isParallelSafe := func(name string) bool {
return name != "" && rt.ToolMeta[name].ReadOnly
}
for i := 0; i < len(toolCalls); {
if ctx.Err() != nil {
fillCancelled(toolCalls[i:])
return results, true
}
// Collect a run of consecutive parallel-safe calls.
j := i + 1
if isParallelSafe(toolCalls[i].Name) {
for j < len(toolCalls) && isParallelSafe(toolCalls[j].Name) {
j++
}
}
batch := toolCalls[i:j]
if len(batch) == 1 {
tr, wasInt := execOne(rt, ctx, batch[0], resultCache)
results = append(results, tr)
if wasInt {
fillCancelled(toolCalls[j:])
return results, true
}
} else {
batchResults, wasInt := execBatch(rt, ctx, batch, resultCache)
results = append(results, batchResults...)
if wasInt {
fillCancelled(toolCalls[j:])
return results, true
}
}
i = j
}
return results, false
}
// execBatch fans out a batch of parallel-safe calls, deduplicating identical ones.
func execBatch(rt *Runtime, ctx TurnCtx, batch []backend.ToolCall, resultCache *sync.Map) ([]toolResult, bool) {
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
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()
uniqueResults[k].tr, uniqueResults[k].wasInt = execOne(rt, ctx, call, resultCache)
}(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,
}
emit(ctx, Event{Role: "call", Name: call.Name, Content: string(call.Arguments)})
emit(ctx, 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 execOne(rt *Runtime, ctx TurnCtx, call backend.ToolCall, resultCache *sync.Map) (toolResult, bool) {
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}
emit(ctx, Event{Role: "tool", Name: call.Name, Content: cr.Content, OutputFormat: toolOutputFormat(rt, call.Name)})
return cr, true
}
emit(ctx, Event{Role: "call", Name: call.Name, Content: string(call.Arguments)})
readSafe := rt.ToolMeta[call.Name].ReadOnly
if readSafe {
key := call.Name + "\x00" + string(call.Arguments)
if v, ok := resultCache.Load(key); ok {
entry := v.(cachedResult)
if cacheValid(entry, call.Arguments) {
emit(ctx, 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
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 ctx.PopInject != nil {
if injected := ctx.PopInject(); injected != "" {
result += "\n\n<system-user-interruption>\n" + injected + "\n</system-user-interruption>"
}
}
emit(ctx, 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 *toolsrv.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 ctx.PopInject != nil {
if injected := ctx.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,
})
}
emit(ctx, 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
}
// ── trackErrors ─────────────────────────────────────────────────────────────
// trackErrors updates consecutive/repeated error tracking and injects nudge
// messages when thresholds are hit. Returns non-nil to abort the loop.
func trackErrors(h *History, ctx TurnCtx, results []toolResult, es *errorState, step int) error {
allErrors := len(results) > 0
allRateLimited := allErrors
for _, r := range results {
if !r.IsError {
allErrors = false
allRateLimited = false
break
}
if !strings.HasPrefix(r.Content, "error: shell is rate-limited") {
allRateLimited = false
}
}
if allRateLimited {
es.consecutive = 0
} else if allErrors {
es.consecutive++
} else {
es.consecutive = 0
}
if es.consecutive >= consecutiveErrorHardLimit {
emit(ctx, Event{Role: "error", Content: fmt.Sprintf("%d consecutive tool errors — aborting", es.consecutive)})
return fmt.Errorf("step %d: %d consecutive tool errors", step, es.consecutive)
}
if es.consecutive == consecutiveErrorSoftLimit {
h.update(backend.Message{
Role: "user",
Content: "<system-consecutive-errors>\nYour last several tool calls all failed. Try a different approach, or ask the user for help. Keep your plan current.\n</system-consecutive-errors>",
}, nil)
}
// Repeated identical error detection.
if allErrors && len(results) > 0 {
var errSig strings.Builder
for _, r := range results {
if r.IsError {
errSig.WriteString(r.Content)
errSig.WriteByte('\n')
}
}
sig := errSig.String()
if sig == es.lastSig {
es.repeatCount++
} else {
es.repeatCount = 1
es.lastSig = sig
}
if es.repeatCount == 3 {
h.update(backend.Message{
Role: "user",
Content: "<system-repeated-error>\nYou have made the same malformed tool call 3 times in a row. Read the error message carefully and fix the arguments. The error was:\n" + strings.TrimSpace(results[0].Content) + "\n</system-repeated-error>",
}, nil)
}
} else {
es.repeatCount = 0
es.lastSig = ""
}
return nil
}
// ── helpers ─────────────────────────────────────────────────────────────────
func emit(ctx TurnCtx, msg Event) {
if ctx.Output != nil {
ctx.Output(msg)
}
}
// contextBudget returns the token threshold (50% of context length) above which
// cold material should be proactively stripped. Returns 0 if unknown.
func contextBudget(ctx context.Context, b backend.Backend) int {
if b == nil {
return 0
}
ctxLen := b.ContextLength(ctx)
if ctxLen <= 0 {
return 0
}
return ctxLen / 2
}
// defaultToolResultMaxBytes caps tool result content sent back to the model.
// This is the single semantic limit for what enters the LLM context window.
//
// Truncation stack (from outer to inner):
// - Server 10MB (toolsrv/shell.go limitedWriter): OOM safety. Prevents runaway
// process output from exhausting memory. Well above model limit.
// - Agent 128KB (this constant): The model-facing cap. Applied in two places:
// streaming path (caps real-time chunks sent to UI) and post-hoc (caps stored
// result). Both use this constant — it's one policy, two enforcement points.
// - Detach 64KB ring (toolsrv/detach.go): Separate system for background
// processes. Unrelated to the model context.
const defaultToolResultMaxBytes = 131072
// MemoryTier classifies how long a tool result stays in the hot message list.
type MemoryTier int
const (
MemoryHot MemoryTier = iota // stays verbatim in messages (default)
MemoryWarm // summarized on next compaction pass
MemoryCold // immediately summarized at update time
)
// toolResult holds the output of a single tool call.
type toolResult struct {
ToolCallID string
Name string
Content string
ContentBlocks []backend.ContentBlock
IsError bool
Tier MemoryTier
}
// transientWait returns the wait duration and whether the error is retryable.
func transientWait(err error, attempt int) (time.Duration, bool) {
var rlErr *backend.RateLimitError
if errors.As(err, &rlErr) {
wait := rlErr.RetryAfter
if wait == 0 {
wait = retryBaseDelay << attempt
}
return wait, true
}
var tErr *backend.TransientError
if errors.As(err, &tErr) {
return time.Duration(2<<attempt) * time.Second, true
}
return 0, false
}
func retryCountdown(ctx TurnCtx, wait time.Duration) error {
deadline := time.Now().Add(wait)
for {
remaining := time.Until(deadline)
if remaining <= 0 {
return nil
}
totalSecs := int(remaining.Seconds()) + 1
h := totalSecs / 3600
m := (totalSecs % 3600) / 60
s := totalSecs % 60
emit(ctx, Event{Role: "retry", Content: fmt.Sprintf("%02d:%02d:%02d", h, m, s)})
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(min(remaining, time.Second)):
}
}
}