1040 lines
32 KiB
Go
1040 lines
32 KiB
Go
package agent
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"ollie/backend"
|
|
"ollie/tools"
|
|
)
|
|
|
|
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
|
|
)
|
|
|
|
// replanGate is the max consecutive tool rounds without a PLAN: block before
|
|
// a mandatory replan nudge is injected.
|
|
const replanGate = 8
|
|
|
|
// stallThreshold is the number of consecutive rounds where LastAction doesn't
|
|
// change before injecting a stall nudge.
|
|
const stallThreshold = 5
|
|
|
|
// planReinjectInterval is the number of tool-call rounds between periodic
|
|
// task state re-injection into the conversation.
|
|
const planReinjectInterval = 10
|
|
|
|
type toolExecutor func(ctx context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error)
|
|
|
|
type agentConfig struct {
|
|
Backend backend.Backend
|
|
Tools []backend.Tool
|
|
Exec toolExecutor
|
|
ClassifyTool func(name string) bool // nil=treat all as serial; true=parallel-read-safe
|
|
ClassifyTier func(name string, args json.RawMessage) ResultTier // nil=TierHot for all; tools self-classify retention
|
|
Output EventHandler
|
|
preamble string // compiled system+agent prompt sent as the system role
|
|
GenerationParams backend.GenerationParams
|
|
PopInject func() string // returns and clears pending inject, or ""
|
|
AutoCompact func(ctx context.Context) // called after each tool round; may compact in-place
|
|
SaveSession func() // called after each state.update(); persists mid-turn progress
|
|
PreTool func(ctx context.Context, name string, args json.RawMessage) HookResult // called before each tool; exit 2 blocks execution
|
|
PostTool func(ctx context.Context, name string, args json.RawMessage, result string) HookResult // called after each tool; exit 0 appends, exit 2 replaces result
|
|
TurnError func(ctx context.Context, errType, errMsg string) HookResult // called on first backend error; if ran, skips retries
|
|
// MaxSteps is the maximum number of tool-call rounds per turn.
|
|
// When reached, a soft nudge is injected and the loop exits cleanly.
|
|
// 0 means unlimited.
|
|
MaxSteps int
|
|
// ReadPlanStep returns the next unchecked step from the plan file, or "".
|
|
// Used as the primary source for TaskState.PlanStep; falls back to
|
|
// heuristic extraction from assistant text if nil.
|
|
ReadPlanStep func() string
|
|
// IncrToolCallCount increments and returns the session-wide tool call counter.
|
|
IncrToolCallCount func() int64
|
|
// ResultCache persists read-safe tool results across turns.
|
|
// Cleared on write operations (file_write, file_edit).
|
|
ResultCache *sync.Map
|
|
}
|
|
|
|
func run(ctx context.Context, cfg agentConfig, state state) error {
|
|
var step int
|
|
var consecutiveErrors int // rounds where every tool call returned an error
|
|
var roundsWithoutPlan int // consecutive tool rounds without a PLAN: block
|
|
var lastAction string // previous round's LastAction for stall detection
|
|
var stallRounds int // consecutive rounds with unchanged LastAction
|
|
var lastErrorSig string // signature of previous round's error(s) for repeat detection
|
|
var repeatErrorCount int // consecutive rounds with identical error signature
|
|
// resultCache stores outputs of read-safe tool calls keyed by name+args.
|
|
// Only tools classified as parallel-safe (immutable reads) are cached.
|
|
// sync.Map is required because execOne may run in concurrent goroutines.
|
|
// Uses the cross-turn cache from agentConfig if available.
|
|
resultCache := cfg.ResultCache
|
|
if resultCache == nil {
|
|
resultCache = &sync.Map{}
|
|
}
|
|
|
|
for {
|
|
emit(cfg, Event{Role: "state", Content: "thinking"})
|
|
// Proactive context gate: strip cold material before calling the backend.
|
|
if budget := contextBudget(ctx, cfg.Backend); budget > 0 && state.estimateTokens() > budget {
|
|
state.stripCold(ctx, cfg.Backend)
|
|
}
|
|
|
|
history := state.history()
|
|
if cfg.preamble != "" {
|
|
history = append([]backend.Message{{Role: "system", Content: cfg.preamble}}, history...)
|
|
}
|
|
if ts := state.taskState(); ts != nil {
|
|
if msg := ts.render(); msg != "" && len(history) > 0 {
|
|
// Insert after system prompt, before conversation.
|
|
history = append([]backend.Message{history[0], {Role: "user", Content: msg}}, history[1:]...)
|
|
}
|
|
}
|
|
|
|
// Stream the assistant's response, retrying on rate limits, transient
|
|
// backend errors (5xx, network), and mid-stream drops. One stable ID is
|
|
// shared by all chunks and the completed assistant message.
|
|
responseID := NewResponseID()
|
|
var content strings.Builder
|
|
var reasoning strings.Builder
|
|
var toolCalls []backend.ToolCall
|
|
var stopReason string
|
|
|
|
for attempt := range maxTransientRetries + 1 {
|
|
content.Reset()
|
|
reasoning.Reset()
|
|
toolCalls = nil
|
|
|
|
ch, err := cfg.Backend.ChatStream(ctx, history, cfg.Tools, cfg.GenerationParams)
|
|
if err != nil {
|
|
if ctx.Err() != nil {
|
|
return ctx.Err()
|
|
}
|
|
// On the first error, fire the turnError hook. If it handles
|
|
// the error (exit 0), return immediately — the hook is
|
|
// responsible for recovery (e.g. switching model and resubmitting).
|
|
if attempt == 0 && cfg.TurnError != nil {
|
|
errType := classifyError(err)
|
|
if r := cfg.TurnError(ctx, errType, err.Error()); r.Handled {
|
|
return fmt.Errorf("step %d: %w", step, err)
|
|
}
|
|
}
|
|
wait, retryable := transientWait(err, attempt)
|
|
if !retryable || attempt >= maxTransientRetries {
|
|
return fmt.Errorf("step %d: %w", step, err)
|
|
}
|
|
var rlErr *backend.RateLimitError
|
|
if errors.As(err, &rlErr) {
|
|
emit(cfg, Event{Role: "limitretry"})
|
|
}
|
|
if err := retryCountdown(ctx, cfg, wait); err != nil {
|
|
return fmt.Errorf("step %d: %w", step, err)
|
|
}
|
|
continue
|
|
}
|
|
|
|
var done bool
|
|
var hadReasoning bool
|
|
for ev := range ch {
|
|
if ev.Reasoning != "" {
|
|
if !hadReasoning {
|
|
emit(cfg, Event{Role: "reasoning", Content: "<think>\n"})
|
|
hadReasoning = true
|
|
}
|
|
reasoning.WriteString(ev.Reasoning)
|
|
emit(cfg, Event{Role: "reasoning", Content: ev.Reasoning})
|
|
}
|
|
if ev.Content != "" {
|
|
if hadReasoning {
|
|
emit(cfg, Event{Role: "reasoning", Content: "\n</think>\n"})
|
|
hadReasoning = false
|
|
}
|
|
content.WriteString(ev.Content)
|
|
emit(cfg, Event{Role: "assistant", Content: ev.Content, ResponseID: responseID})
|
|
}
|
|
toolCalls = append(toolCalls, ev.ToolCalls...)
|
|
if ev.Done {
|
|
if hadReasoning {
|
|
emit(cfg, Event{Role: "reasoning", Content: "\n</think>\n"})
|
|
hadReasoning = false
|
|
}
|
|
stopReason = ev.StopReason
|
|
done = true
|
|
if ev.Usage.InputTokens > 0 || ev.Usage.OutputTokens > 0 {
|
|
emit(cfg, 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 {
|
|
// Backend didn't report usage; estimate from content.
|
|
inChars := 0
|
|
for _, m := range history {
|
|
inChars += len(m.Content)
|
|
for _, tc := range m.ToolCalls {
|
|
inChars += len(tc.Name) + len(tc.Arguments)
|
|
}
|
|
}
|
|
emit(cfg, Event{
|
|
Role: "usage",
|
|
Content: fmt.Sprintf("%d %d 1 0", inChars/4, content.Len()/4),
|
|
})
|
|
}
|
|
break
|
|
}
|
|
}
|
|
|
|
if done {
|
|
break
|
|
}
|
|
|
|
if ctx.Err() != nil {
|
|
// User pause: record partial state and return.
|
|
if hadReasoning {
|
|
emit(cfg, Event{Role: "reasoning", Content: "\n</think>\n"})
|
|
}
|
|
msg := backend.Message{ID: responseID, Role: "assistant", Content: content.String(), Reasoning: reasoning.String(), ToolCalls: toolCalls}
|
|
var results []toolResult
|
|
for _, tc := range toolCalls {
|
|
results = append(results, toolResult{
|
|
ToolCallID: tc.ID,
|
|
Name: tc.Name,
|
|
Content: `{"status":"cancelled","error":"stream interrupted"}`,
|
|
IsError: true,
|
|
})
|
|
}
|
|
state.update(msg, results)
|
|
if cfg.SaveSession != nil {
|
|
cfg.SaveSession()
|
|
}
|
|
return ctx.Err()
|
|
}
|
|
|
|
// Pure stream drop — retry if attempts remain.
|
|
if attempt >= maxTransientRetries {
|
|
return fmt.Errorf("step %d: stream dropped (no more retries)", step)
|
|
}
|
|
if hadReasoning {
|
|
emit(cfg, Event{Role: "reasoning", Content: "\n</think>\n"})
|
|
}
|
|
wait := streamDropBaseDelay << attempt
|
|
if err := retryCountdown(ctx, cfg, wait); err != nil {
|
|
return fmt.Errorf("step %d: %w", step, err)
|
|
}
|
|
}
|
|
|
|
switch stopReason {
|
|
case "stop", "tool_calls", "length", "error", "":
|
|
// normal
|
|
default:
|
|
return fmt.Errorf("step %d: %s", step, stopReason)
|
|
}
|
|
|
|
// Parse text-based tool calls: models that don't support the function
|
|
// calling API emit tool invocations as plain text (e.g.
|
|
// "file_read: args=[...]" (text-based tool call syntax)). Parse these into proper ToolCall structs.
|
|
if len(toolCalls) == 0 && len(cfg.Tools) > 0 {
|
|
if parsed := parseTextToolCalls(content.String(), cfg.Tools); len(parsed) > 0 {
|
|
toolCalls = parsed
|
|
}
|
|
}
|
|
|
|
// Execute tool calls, running consecutive parallel-read-safe tools concurrently.
|
|
msg := backend.Message{ID: responseID, Role: "assistant", Content: content.String(), Reasoning: reasoning.String(), ToolCalls: toolCalls}
|
|
results := make([]toolResult, 0, len(toolCalls))
|
|
interrupted := false
|
|
|
|
cancelledResult := func(tc backend.ToolCall) toolResult {
|
|
return toolResult{
|
|
ToolCallID: tc.ID,
|
|
Name: tc.Name,
|
|
Content: `{"status":"cancelled","error":"interrupted"}`,
|
|
IsError: true,
|
|
}
|
|
}
|
|
|
|
// execOne runs a single tool call end-to-end and reports whether the
|
|
// context was cancelled during execution.
|
|
execOne := func(tc backend.ToolCall) (toolResult, bool) {
|
|
if tc.Name == "" {
|
|
return toolResult{ToolCallID: tc.ID, Name: tc.Name, Content: "error: empty tool name", IsError: true}, false
|
|
}
|
|
if ctx.Err() != nil {
|
|
cr := cancelledResult(tc)
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: cr.Content})
|
|
return cr, true
|
|
}
|
|
emit(cfg, Event{Role: "call", Name: tc.Name, Content: string(tc.Arguments)})
|
|
if cfg.PreTool != nil {
|
|
hr := cfg.PreTool(ctx, tc.Name, tc.Arguments)
|
|
if hr.Failed > 0 {
|
|
emit(cfg, Event{Role: "info", Content: "preTool: " + hr.Summary()})
|
|
}
|
|
if hr.Blocked {
|
|
blocked := hr.Context
|
|
if blocked == "" {
|
|
blocked = fmt.Sprintf("tool %q blocked by hook", tc.Name)
|
|
}
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: blocked})
|
|
return toolResult{ToolCallID: tc.ID, Name: tc.Name, Content: blocked, IsError: true}, false
|
|
}
|
|
}
|
|
readSafe := cfg.ClassifyTool != nil && cfg.ClassifyTool(tc.Name)
|
|
if readSafe {
|
|
key := tc.Name + "\x00" + string(tc.Arguments)
|
|
if v, ok := resultCache.Load(key); ok {
|
|
cached := v.(string)
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: cached})
|
|
return toolResult{ToolCallID: tc.ID, Name: tc.Name, Content: cached}, false
|
|
}
|
|
}
|
|
var result string
|
|
var resultBlocks []backend.ContentBlock
|
|
var isErr bool
|
|
streamed := false
|
|
if cfg.Exec != nil {
|
|
streamBytes := 0
|
|
streamCtx := tools.WithOutputStream(ctx, func(data string) {
|
|
if streamBytes >= defaultToolResultMaxBytes {
|
|
return // already at ceiling, drop further chunks
|
|
}
|
|
if !streamed {
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name})
|
|
streamed = true
|
|
}
|
|
streamBytes += len(data)
|
|
if streamBytes > defaultToolResultMaxBytes {
|
|
// Emit only the portion within the ceiling.
|
|
excess := streamBytes - defaultToolResultMaxBytes
|
|
data = data[:len(data)-excess]
|
|
}
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: data})
|
|
})
|
|
out, blocks, err := cfg.Exec(streamCtx, tc.Name, tc.Arguments)
|
|
if err != nil {
|
|
isErr = true
|
|
if ctx.Err() != nil {
|
|
result = "error: tool execution interrupted by user"
|
|
if cfg.PopInject != nil {
|
|
if injected := cfg.PopInject(); injected != "" {
|
|
result += "\n\n<system-user-interruption>\n" + injected + "\n</system-user-interruption>"
|
|
}
|
|
}
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: result})
|
|
return toolResult{ToolCallID: tc.ID, Name: tc.Name, Content: result, IsError: true}, true
|
|
}
|
|
result = fmt.Sprintf("error: %v", err)
|
|
} else {
|
|
result = out
|
|
resultBlocks = blocks
|
|
}
|
|
} else {
|
|
result = "error: no tool executor configured"
|
|
isErr = true
|
|
}
|
|
// Accumulate suffix text (PostTool context, user-interruptions,
|
|
// truncation hints) that must be emitted after streaming completes.
|
|
var suffix string
|
|
if cfg.PostTool != nil {
|
|
hr := cfg.PostTool(ctx, tc.Name, tc.Arguments, result)
|
|
if hr.Failed > 0 {
|
|
emit(cfg, Event{Role: "info", Content: "postTool: " + hr.Summary()})
|
|
}
|
|
if hr.Blocked {
|
|
result = hr.Context
|
|
} else if hr.Context != "" {
|
|
suffix += "\n" + hr.Context
|
|
}
|
|
}
|
|
if cfg.PopInject != nil {
|
|
if injected := cfg.PopInject(); injected != "" {
|
|
suffix += "\n\n<system-user-interruption>\n" + injected + "\n</system-user-interruption>"
|
|
}
|
|
}
|
|
result += suffix
|
|
// Safety ceiling: unconditionally cap all tool results at 128KB.
|
|
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.]",
|
|
tc.Name, defaultToolResultMaxBytes, orig)
|
|
result += suffix
|
|
}
|
|
if readSafe && !isErr {
|
|
resultCache.Store(tc.Name+"\x00"+string(tc.Arguments), result)
|
|
}
|
|
// Invalidate cache on write operations that may change file contents.
|
|
if !readSafe && (tc.Name == "file_write" || tc.Name == "file_edit") {
|
|
*resultCache = sync.Map{}
|
|
}
|
|
if !streamed {
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: result})
|
|
} else if suffix != "" {
|
|
// Emit suffixes that were appended after streaming completed
|
|
// so they appear at the end of the chat output.
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: suffix})
|
|
}
|
|
tier := TierHot
|
|
if !isErr && cfg.ClassifyTier != nil {
|
|
tier = cfg.ClassifyTier(tc.Name, tc.Arguments)
|
|
}
|
|
return toolResult{ToolCallID: tc.ID, Name: tc.Name, Content: result, ContentBlocks: resultBlocks, IsError: isErr, Tier: tier}, false
|
|
}
|
|
|
|
isParallelSafe := func(name string) bool {
|
|
return name != "" && cfg.ClassifyTool != nil && cfg.ClassifyTool(name)
|
|
}
|
|
|
|
for i := 0; i < len(toolCalls) && !interrupted; {
|
|
if ctx.Err() != nil {
|
|
for _, remaining := range toolCalls[i:] {
|
|
cr := cancelledResult(remaining)
|
|
emit(cfg, Event{Role: "tool", Name: remaining.Name, Content: cr.Content})
|
|
results = append(results, cr)
|
|
}
|
|
interrupted = true
|
|
break
|
|
}
|
|
|
|
// 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]
|
|
|
|
fillCancelled := func(from int) {
|
|
for _, remaining := range toolCalls[from:] {
|
|
cr := cancelledResult(remaining)
|
|
emit(cfg, Event{Role: "tool", Name: remaining.Name, Content: cr.Content})
|
|
results = append(results, cr)
|
|
}
|
|
}
|
|
|
|
if len(batch) == 1 {
|
|
tr, wasInt := execOne(batch[0])
|
|
results = append(results, tr)
|
|
if wasInt {
|
|
fillCancelled(j)
|
|
interrupted = true
|
|
}
|
|
} else {
|
|
// Fan out the batch concurrently; preserve submission order in results.
|
|
// Deduplicate identical calls so only one executes per unique key.
|
|
type inflightResult struct {
|
|
tr toolResult
|
|
wasInt bool
|
|
}
|
|
inflight := make(map[string]int) // key -> index of first occurrence
|
|
batchResults := make([]toolResult, len(batch))
|
|
batchInt := make([]bool, len(batch))
|
|
var wg sync.WaitGroup
|
|
uniqueResults := make([]inflightResult, len(batch))
|
|
for k, tc := range batch {
|
|
key := tc.Name + "\x00" + string(tc.Arguments)
|
|
if first, dup := inflight[key]; dup {
|
|
// Will copy result from first occurrence after wg.Wait.
|
|
batchResults[k] = toolResult{} // placeholder
|
|
_ = first // used below
|
|
continue
|
|
}
|
|
inflight[key] = k
|
|
wg.Add(1)
|
|
go func(k int, tc backend.ToolCall) {
|
|
defer wg.Done()
|
|
uniqueResults[k].tr, uniqueResults[k].wasInt = execOne(tc)
|
|
}(k, tc)
|
|
}
|
|
wg.Wait()
|
|
// Fill results: unique calls get their own result, duplicates copy from first.
|
|
for k, tc := range batch {
|
|
key := tc.Name + "\x00" + string(tc.Arguments)
|
|
first := inflight[key]
|
|
if k == first {
|
|
batchResults[k] = uniqueResults[k].tr
|
|
batchInt[k] = uniqueResults[k].wasInt
|
|
} else {
|
|
batchResults[k] = toolResult{
|
|
ToolCallID: tc.ID,
|
|
Name: tc.Name,
|
|
Content: uniqueResults[first].tr.Content,
|
|
IsError: uniqueResults[first].tr.IsError,
|
|
}
|
|
batchInt[k] = uniqueResults[first].wasInt
|
|
emit(cfg, Event{Role: "call", Name: tc.Name, Content: string(tc.Arguments)})
|
|
emit(cfg, Event{Role: "tool", Name: tc.Name, Content: batchResults[k].Content})
|
|
}
|
|
}
|
|
// Add all batch results — model needs a result for every tool call.
|
|
results = append(results, batchResults...)
|
|
for _, wasInt := range batchInt {
|
|
if wasInt {
|
|
fillCancelled(j)
|
|
interrupted = true
|
|
break
|
|
}
|
|
}
|
|
}
|
|
i = j
|
|
}
|
|
|
|
state.update(msg, results)
|
|
if cfg.SaveSession != nil {
|
|
cfg.SaveSession()
|
|
}
|
|
|
|
// Auto-update TaskState from what actually happened this round.
|
|
if ts := state.taskState(); ts != nil {
|
|
inferTaskStateUpdate(ts, msg, results, cfg.ReadPlanStep)
|
|
state.updateTaskState(*ts)
|
|
}
|
|
|
|
if interrupted {
|
|
return ctx.Err()
|
|
}
|
|
|
|
if len(toolCalls) == 0 {
|
|
break
|
|
}
|
|
|
|
// Track consecutive rounds where every tool call errored.
|
|
allErrors := len(results) > 0
|
|
for _, r := range results {
|
|
if !r.IsError {
|
|
allErrors = false
|
|
break
|
|
}
|
|
}
|
|
if allErrors {
|
|
consecutiveErrors++
|
|
} else {
|
|
consecutiveErrors = 0
|
|
}
|
|
if consecutiveErrors >= consecutiveErrorHardLimit {
|
|
emit(cfg, Event{Role: "error", Content: fmt.Sprintf("%d consecutive tool errors — aborting", consecutiveErrors)})
|
|
return fmt.Errorf("step %d: %d consecutive tool errors", step, consecutiveErrors)
|
|
}
|
|
if consecutiveErrors == consecutiveErrorSoftLimit {
|
|
// Nudge the model to try a different approach and keep the plan current.
|
|
state.update(backend.Message{
|
|
Role: "user",
|
|
Content: "[system: your last several tool calls all failed. Try a different approach, or ask the user for help. Keep your plan current — update s/$OLLIE_SESSION_ID/plan to reflect where you are and what remains.]",
|
|
}, nil)
|
|
}
|
|
|
|
// Repeated identical error detection: if the same error occurs 3 times
|
|
// in a row, the model is stuck on a format/argument issue. Inject the
|
|
// error content as a corrective nudge immediately rather than waiting
|
|
// for the generic soft limit.
|
|
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 == lastErrorSig {
|
|
repeatErrorCount++
|
|
} else {
|
|
repeatErrorCount = 1
|
|
lastErrorSig = sig
|
|
}
|
|
if repeatErrorCount == 3 {
|
|
state.update(backend.Message{
|
|
Role: "user",
|
|
Content: "[system: you have made the same malformed tool call 3 times in a row. Read the error message carefully and fix the arguments. The error was: " + strings.TrimSpace(results[0].Content) + "]",
|
|
}, nil)
|
|
}
|
|
} else {
|
|
repeatErrorCount = 0
|
|
lastErrorSig = ""
|
|
}
|
|
|
|
// Replan gate: if the model has been calling tools without emitting a
|
|
// PLAN: block, force it to replan before continuing.
|
|
if strings.Contains(msg.Content, "PLAN:") {
|
|
roundsWithoutPlan = 0
|
|
} else {
|
|
roundsWithoutPlan++
|
|
}
|
|
if roundsWithoutPlan >= replanGate {
|
|
state.update(backend.Message{
|
|
Role: "user",
|
|
Content: "[system: you have executed " + fmt.Sprintf("%d", replanGate) + " tool rounds without replanning. Stop and write a PLAN: block showing your current checklist before calling any more tools. Update s/$OLLIE_SESSION_ID/plan.]",
|
|
}, nil)
|
|
roundsWithoutPlan = 0
|
|
}
|
|
|
|
// Stall detection: if LastAction hasn't changed across rounds, the
|
|
// model is likely stuck in a loop.
|
|
if ts := state.taskState(); ts != nil {
|
|
if ts.LastAction == lastAction && lastAction != "" {
|
|
stallRounds++
|
|
} else {
|
|
stallRounds = 0
|
|
}
|
|
lastAction = ts.LastAction
|
|
if stallRounds >= stallThreshold {
|
|
state.update(backend.Message{
|
|
Role: "user",
|
|
Content: "[system: stall detected — you have repeated the same action for " + fmt.Sprintf("%d", stallRounds) + " rounds. Step back, reassess your approach, and try something different. Update your plan.]",
|
|
}, nil)
|
|
stallRounds = 0
|
|
}
|
|
}
|
|
|
|
// Soft step-budget guardrail: when MaxSteps is set and the budget is
|
|
// exhausted, nudge the model to wrap up and exit cleanly. This is not
|
|
// a hard abort — the model gets one final turn without tools to emit
|
|
// a summary or hand-off message.
|
|
if cfg.MaxSteps > 0 && step >= cfg.MaxSteps-1 {
|
|
emit(cfg, Event{Role: "maxsteps", Content: fmt.Sprintf("%d", step+1)})
|
|
state.update(backend.Message{
|
|
Role: "user",
|
|
Content: fmt.Sprintf("[system: step budget exhausted (%d/%d steps used). Stop calling tools. Summarize what you have done and what remains, then stop.]", step+1, cfg.MaxSteps),
|
|
}, nil)
|
|
break
|
|
}
|
|
|
|
// Periodic task state re-injection: every N tool rounds, surface the
|
|
// structured task state back into the conversation to keep the model on track.
|
|
if step > 0 && step%planReinjectInterval == 0 {
|
|
if ts := state.taskState(); ts != nil {
|
|
if msg := ts.render(); msg != "" {
|
|
state.update(backend.Message{
|
|
Role: "user",
|
|
Content: "[system: review your current task state and continue. Update it if your approach has changed.]\n\n" + msg,
|
|
}, nil)
|
|
}
|
|
}
|
|
}
|
|
|
|
if cfg.AutoCompact != nil {
|
|
cfg.AutoCompact(ctx)
|
|
}
|
|
|
|
step++
|
|
}
|
|
|
|
return nil
|
|
}
|
|
|
|
func emit(cfg agentConfig, msg Event) {
|
|
if cfg.Output != nil {
|
|
cfg.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
|
|
}
|
|
|
|
// inferTaskStateUpdate updates TaskState fields from the actual tool round
|
|
// without requiring model cooperation. Keeps the state current even if the
|
|
// model ignores re-injection nudges.
|
|
func inferTaskStateUpdate(ts *TaskState, msg backend.Message, results []toolResult, readPlanStep func() string) {
|
|
// LastAction: summarize what tools ran and whether they succeeded.
|
|
if len(msg.ToolCalls) > 0 {
|
|
names := make([]string, 0, len(msg.ToolCalls))
|
|
for _, tc := range msg.ToolCalls {
|
|
names = append(names, tc.Name)
|
|
}
|
|
errCount := 0
|
|
for _, r := range results {
|
|
if r.IsError {
|
|
errCount++
|
|
}
|
|
}
|
|
action := strings.Join(names, ", ")
|
|
if errCount > 0 {
|
|
action += fmt.Sprintf(" (%d/%d failed)", errCount, len(results))
|
|
}
|
|
ts.LastAction = action
|
|
}
|
|
|
|
// PlanStep: prefer the plan file (ground truth), fall back to text heuristic.
|
|
if readPlanStep != nil {
|
|
if step := readPlanStep(); step != "" {
|
|
ts.PlanStep = step
|
|
return
|
|
}
|
|
}
|
|
// Fallback: extract from PLAN: block if present in assistant text.
|
|
if idx := strings.Index(msg.Content, "PLAN:"); idx >= 0 {
|
|
lines := strings.Split(msg.Content[idx:], "\n")
|
|
for _, line := range lines[1:] {
|
|
trimmed := strings.TrimSpace(line)
|
|
if strings.HasPrefix(trimmed, "- [ ]") {
|
|
ts.PlanStep = strings.TrimSpace(trimmed[5:])
|
|
break
|
|
}
|
|
if trimmed == "" || (!strings.HasPrefix(trimmed, "- [") && trimmed != "") {
|
|
break
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
// parseTextToolCalls extracts tool calls from assistant text when the model
|
|
// emits them as plain text instead of using the function calling API.
|
|
// Recognizes formats:
|
|
// - tool_name: {"key": ...} (JSON object)
|
|
// - tool_name: key=[...], key2=... (shorthand: key=value pairs → JSON object)
|
|
//
|
|
// Returns nil if no valid tool calls are found.
|
|
func parseTextToolCalls(text string, tools []backend.Tool) []backend.ToolCall {
|
|
toolNames := make(map[string]bool, len(tools))
|
|
for _, t := range tools {
|
|
toolNames[t.Name] = true
|
|
}
|
|
|
|
var calls []backend.ToolCall
|
|
for _, name := range []string{"shell"} {
|
|
if !toolNames[name] {
|
|
continue
|
|
}
|
|
prefix := name + ":"
|
|
idx := strings.Index(text, prefix)
|
|
if idx < 0 {
|
|
continue
|
|
}
|
|
rest := strings.TrimSpace(text[idx+len(prefix):])
|
|
if rest == "" {
|
|
continue
|
|
}
|
|
|
|
var argsJSON string
|
|
if rest[0] == '{' {
|
|
// Direct JSON object
|
|
argsJSON = extractJSONObject(rest)
|
|
if argsJSON == "" {
|
|
argsJSON = extractJSONObject(relaxJSON(rest))
|
|
}
|
|
} else {
|
|
// Shorthand format: key=[...], key2=value
|
|
argsJSON = shorthandToJSON(rest)
|
|
}
|
|
if argsJSON == "" {
|
|
continue
|
|
}
|
|
var check json.RawMessage
|
|
if json.Unmarshal([]byte(argsJSON), &check) != nil {
|
|
continue
|
|
}
|
|
calls = append(calls, backend.ToolCall{
|
|
ID: fmt.Sprintf("text-%s-%d", name, len(calls)),
|
|
Name: name,
|
|
Arguments: json.RawMessage(argsJSON),
|
|
})
|
|
}
|
|
return calls
|
|
}
|
|
|
|
// shorthandToJSON converts "key=[...], key2=value" format to a JSON object.
|
|
// Handles: steps=[{code: "date"}], timeout=30
|
|
func shorthandToJSON(s string) string {
|
|
// Find the first key=value where value starts with [ or {
|
|
eqIdx := strings.IndexByte(s, '=')
|
|
if eqIdx < 0 {
|
|
return ""
|
|
}
|
|
key := strings.TrimSpace(s[:eqIdx])
|
|
if key == "" {
|
|
return ""
|
|
}
|
|
valStr := s[eqIdx+1:]
|
|
|
|
// Extract the value (balanced brackets/braces)
|
|
val := extractBalanced(strings.TrimSpace(valStr))
|
|
if val == "" {
|
|
return ""
|
|
}
|
|
|
|
// Relax the value's JSON (unquoted keys → quoted)
|
|
relaxed := relaxJSON(val)
|
|
|
|
// Build the JSON object: {"key": relaxed_value}
|
|
keyJSON, _ := json.Marshal(key)
|
|
result := "{" + string(keyJSON) + ":" + relaxed + "}"
|
|
|
|
// Check if there are more key=value pairs after the value
|
|
afterVal := strings.TrimSpace(valStr[len(val):])
|
|
if len(afterVal) > 0 && afterVal[0] == ',' {
|
|
// Parse additional key=value pairs
|
|
extra := parseExtraKV(afterVal[1:])
|
|
if extra != "" {
|
|
// Merge: strip trailing } from result, append extra
|
|
result = result[:len(result)-1] + "," + extra + "}"
|
|
}
|
|
}
|
|
|
|
return result
|
|
}
|
|
|
|
// extractBalanced extracts a balanced [...] or {...} from the start of s.
|
|
func extractBalanced(s string) string {
|
|
if len(s) == 0 {
|
|
return ""
|
|
}
|
|
open := s[0]
|
|
var close byte
|
|
switch open {
|
|
case '[':
|
|
close = ']'
|
|
case '{':
|
|
close = '}'
|
|
default:
|
|
// Simple value (number, string) — take until comma or newline
|
|
end := strings.IndexAny(s, ",\n")
|
|
if end < 0 {
|
|
return strings.TrimSpace(s)
|
|
}
|
|
return strings.TrimSpace(s[:end])
|
|
}
|
|
|
|
depth := 0
|
|
inStr := false
|
|
escaped := false
|
|
for i := 0; i < len(s); i++ {
|
|
c := s[i]
|
|
if escaped {
|
|
escaped = false
|
|
continue
|
|
}
|
|
if c == '\\' && inStr {
|
|
escaped = true
|
|
continue
|
|
}
|
|
if c == '"' {
|
|
inStr = !inStr
|
|
continue
|
|
}
|
|
if inStr {
|
|
continue
|
|
}
|
|
if c == open {
|
|
depth++
|
|
} else if c == close {
|
|
depth--
|
|
if depth == 0 {
|
|
return s[:i+1]
|
|
}
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
// parseExtraKV parses "key=value, key2=value2" into JSON fields (without outer braces).
|
|
func parseExtraKV(s string) string {
|
|
s = strings.TrimSpace(s)
|
|
if s == "" {
|
|
return ""
|
|
}
|
|
var parts []string
|
|
for s != "" {
|
|
eqIdx := strings.IndexByte(s, '=')
|
|
if eqIdx < 0 {
|
|
break
|
|
}
|
|
key := strings.TrimSpace(s[:eqIdx])
|
|
s = strings.TrimSpace(s[eqIdx+1:])
|
|
val := extractBalanced(s)
|
|
if val == "" {
|
|
break
|
|
}
|
|
s = strings.TrimSpace(s[len(val):])
|
|
if len(s) > 0 && s[0] == ',' {
|
|
s = s[1:]
|
|
}
|
|
keyJSON, _ := json.Marshal(key)
|
|
parts = append(parts, string(keyJSON)+":"+relaxJSON(val))
|
|
}
|
|
return strings.Join(parts, ",")
|
|
}
|
|
|
|
// extractJSONObject finds the first balanced {...} in s.
|
|
func extractJSONObject(s string) string {
|
|
start := strings.IndexByte(s, '{')
|
|
if start < 0 {
|
|
return ""
|
|
}
|
|
depth := 0
|
|
inStr := false
|
|
escaped := false
|
|
for i := start; i < len(s); i++ {
|
|
c := s[i]
|
|
if escaped {
|
|
escaped = false
|
|
continue
|
|
}
|
|
if c == '\\' && inStr {
|
|
escaped = true
|
|
continue
|
|
}
|
|
if c == '"' {
|
|
inStr = !inStr
|
|
continue
|
|
}
|
|
if inStr {
|
|
continue
|
|
}
|
|
if c == '{' {
|
|
depth++
|
|
} else if c == '}' {
|
|
depth--
|
|
if depth == 0 {
|
|
return s[start : i+1]
|
|
}
|
|
}
|
|
}
|
|
return ""
|
|
}
|
|
|
|
// relaxJSON converts common relaxed-JSON patterns to valid JSON:
|
|
// - unquoted keys (word: → "word":)
|
|
// - single-quoted strings → double-quoted
|
|
// This is a best-effort heuristic for model output.
|
|
func relaxJSON(s string) string {
|
|
var out strings.Builder
|
|
out.Grow(len(s))
|
|
i := 0
|
|
for i < len(s) {
|
|
c := s[i]
|
|
// Single-quoted string → double-quoted
|
|
if c == '\'' {
|
|
out.WriteByte('"')
|
|
i++
|
|
for i < len(s) && s[i] != '\'' {
|
|
if s[i] == '"' {
|
|
out.WriteString(`\"`)
|
|
} else if s[i] == '\\' && i+1 < len(s) {
|
|
out.WriteByte(s[i])
|
|
i++
|
|
out.WriteByte(s[i])
|
|
} else {
|
|
out.WriteByte(s[i])
|
|
}
|
|
i++
|
|
}
|
|
out.WriteByte('"')
|
|
if i < len(s) {
|
|
i++ // skip closing '
|
|
}
|
|
continue
|
|
}
|
|
// Unquoted key before colon: word followed by optional whitespace then ':'
|
|
if c >= 'a' && c <= 'z' || c >= 'A' && c <= 'Z' || c == '_' {
|
|
j := i
|
|
for j < len(s) && (s[j] >= 'a' && s[j] <= 'z' || s[j] >= 'A' && s[j] <= 'Z' || s[j] >= '0' && s[j] <= '9' || s[j] == '_') {
|
|
j++
|
|
}
|
|
// Check if followed by optional whitespace then ':'
|
|
k := j
|
|
for k < len(s) && (s[k] == ' ' || s[k] == '\t') {
|
|
k++
|
|
}
|
|
if k < len(s) && s[k] == ':' {
|
|
// It's an unquoted key — quote it
|
|
out.WriteByte('"')
|
|
out.WriteString(s[i:j])
|
|
out.WriteByte('"')
|
|
i = j
|
|
continue
|
|
}
|
|
// Not a key, just copy the word
|
|
out.WriteString(s[i:j])
|
|
i = j
|
|
continue
|
|
}
|
|
out.WriteByte(c)
|
|
i++
|
|
}
|
|
return out.String()
|
|
}
|
|
|
|
// classifyError returns a short string identifying the error type for the
|
|
// turnError hook payload.
|
|
func classifyError(err error) string {
|
|
var rlErr *backend.RateLimitError
|
|
if errors.As(err, &rlErr) {
|
|
return "rate_limit"
|
|
}
|
|
var tuErr *backend.ToolUnsupportedError
|
|
if errors.As(err, &tuErr) {
|
|
return "tool_unsupported"
|
|
}
|
|
var coErr *backend.ContextOverflowError
|
|
if errors.As(err, &coErr) {
|
|
return "context_overflow"
|
|
}
|
|
var tErr *backend.TransientError
|
|
if errors.As(err, &tErr) {
|
|
return "transient"
|
|
}
|
|
return "unknown"
|
|
}
|
|
|
|
// transientWait returns the retry wait for a retryable error and whether it is
|
|
// retryable. Rate limits use longer waits; transient/network errors use shorter.
|
|
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 context.Context, cfg agentConfig, wait time.Duration) error {
|
|
deadline := time.Now().Add(wait)
|
|
for {
|
|
remaining := time.Until(deadline)
|
|
if remaining <= 0 {
|
|
return nil
|
|
}
|
|
secs := int(remaining.Seconds()) + 1
|
|
emit(cfg, Event{Role: "retry", Content: fmt.Sprintf("%d", secs)})
|
|
select {
|
|
case <-ctx.Done():
|
|
return ctx.Err()
|
|
case <-time.After(min(remaining, time.Second)):
|
|
}
|
|
}
|
|
}
|