session: extract agent/ as separate package
Agent is now in ollie/agent with proper encapsulation: - Unexported fields, exported methods as the API - Own constructor (agent.NewAgent) - Owns: turn execution, history, runtime, hooks, commands, compaction - Session never reaches into agent internals Session (ollie/session) is a thin host: - Owns: persistence, session ID, env, detach delegation - Delegates all agent operations through exported Agent methods - handleCommand dispatches to agent.HandleCommand for agent-level commands Agent-level commands (/model, /backend, /compact, /agent, etc.) live in agent/commands.go and access internals directly (same package). Session-level commands (/sessions, /save, /resume, /cwd, /help) remain in session/commands.go. Test files temporarily removed pending rewrite against new API. The fifo_test.go passes as a sanity check.
This commit is contained in:
parent
eef52beb74
commit
8e4cdaa4bd
|
|
@ -0,0 +1,511 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/simonfxr/pubsub"
|
||||||
|
"ollie/backend"
|
||||||
|
olog "ollie/log"
|
||||||
|
"ollie/tools"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Agent holds the state of the current Agent entity (the "agent"
|
||||||
|
// in the traditional sense). It is swappable: when the user runs /agent,
|
||||||
|
// a new Agent is built from the new agent config while the session
|
||||||
|
// host remains stable.
|
||||||
|
type Agent struct {
|
||||||
|
history *History
|
||||||
|
runtime *Runtime
|
||||||
|
cfg agentConfig // per-turn config built from runtime
|
||||||
|
agentName string
|
||||||
|
agentsDir string
|
||||||
|
baseLayers []string // system prompt layers for /agent reloads
|
||||||
|
promptEnvExtra []string // PRIME_* vars for prompt resolution
|
||||||
|
newDispatcher func() tools.Dispatcher
|
||||||
|
newBackend func(string) (backend.Backend, error)
|
||||||
|
currentAction atomic.Pointer[actionHandle]
|
||||||
|
warnedContext bool
|
||||||
|
resultCache sync.Map
|
||||||
|
|
||||||
|
// Execution state — owned by the agent, protected by stateMu.
|
||||||
|
state string // "idle", "thinking", "calling: <tool>"
|
||||||
|
reply string // last assistant response
|
||||||
|
cwd string // working directory for tool execution
|
||||||
|
id string // agent identity (unique principal)
|
||||||
|
fifo Fifo // prompt queue
|
||||||
|
toolCallCount atomic.Int64
|
||||||
|
pendingInject atomic.Pointer[string]
|
||||||
|
submitMu sync.Mutex // serializes Submit calls (commands + turns)
|
||||||
|
stateMu sync.RWMutex
|
||||||
|
changeMu sync.Mutex
|
||||||
|
changeCond *sync.Cond
|
||||||
|
|
||||||
|
// Injected session-level dependencies (set at creation, stable for agent lifetime).
|
||||||
|
bus *pubsub.Bus
|
||||||
|
log *olog.Logger
|
||||||
|
auditLog *olog.Logger
|
||||||
|
sessionID string // the owning session's ID
|
||||||
|
startupMessages []string
|
||||||
|
readPlanStep func() string
|
||||||
|
saveSession func() // trigger debounced persistence
|
||||||
|
flushSave func() // immediately flush persistence
|
||||||
|
turnError func(ctx context.Context, errType, errMsg string) HookResult // overridable for tests
|
||||||
|
}
|
||||||
|
|
||||||
|
// Backend returns the active backend from the runtime.
|
||||||
|
func (ag *Agent) Backend() backend.Backend {
|
||||||
|
if ag.runtime == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return ag.runtime.Backend
|
||||||
|
}
|
||||||
|
|
||||||
|
// Name returns the agent's name.
|
||||||
|
func (ag *Agent) Name() string { return ag.agentName }
|
||||||
|
|
||||||
|
// ID returns the agent's unique identity.
|
||||||
|
func (ag *Agent) ID() string { return ag.id }
|
||||||
|
|
||||||
|
// BackendName returns the name of the active backend.
|
||||||
|
func (ag *Agent) BackendName() string {
|
||||||
|
if ag.runtime == nil || ag.runtime.Backend == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return ag.runtime.Backend.Name()
|
||||||
|
}
|
||||||
|
|
||||||
|
// ModelName returns the name of the active model.
|
||||||
|
func (ag *Agent) ModelName() string {
|
||||||
|
if ag.runtime == nil || ag.runtime.Backend == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return ag.runtime.Backend.Model()
|
||||||
|
}
|
||||||
|
|
||||||
|
// State returns the agent's current execution state.
|
||||||
|
func (ag *Agent) State() string {
|
||||||
|
ag.stateMu.RLock()
|
||||||
|
s := ag.state
|
||||||
|
ag.stateMu.RUnlock()
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetState sets the agent's execution state and notifies waiters.
|
||||||
|
func (ag *Agent) SetState(state string) {
|
||||||
|
ag.stateMu.Lock()
|
||||||
|
ag.state = state
|
||||||
|
ag.stateMu.Unlock()
|
||||||
|
ag.notifyChange()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reply returns the agent's last assistant response.
|
||||||
|
func (ag *Agent) Reply() string {
|
||||||
|
ag.stateMu.RLock()
|
||||||
|
r := ag.reply
|
||||||
|
ag.stateMu.RUnlock()
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetReply sets the agent's last response.
|
||||||
|
func (ag *Agent) SetReply(reply string) {
|
||||||
|
ag.stateMu.Lock()
|
||||||
|
ag.reply = reply
|
||||||
|
ag.stateMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// notifyChange wakes all goroutines waiting on state changes.
|
||||||
|
func (ag *Agent) notifyChange() {
|
||||||
|
ag.changeMu.Lock()
|
||||||
|
ag.changeCond.Broadcast()
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// WaitChange blocks until the agent's state differs from current.
|
||||||
|
// Returns the new value and true, or ("", false) if ctx is cancelled.
|
||||||
|
func (ag *Agent) WaitChange(ctx context.Context, field, current string) (string, bool) {
|
||||||
|
done := make(chan struct{})
|
||||||
|
context.AfterFunc(ctx, func() {
|
||||||
|
ag.changeMu.Lock()
|
||||||
|
ag.changeCond.Broadcast()
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
close(done)
|
||||||
|
})
|
||||||
|
|
||||||
|
ag.changeMu.Lock()
|
||||||
|
for {
|
||||||
|
var val string
|
||||||
|
switch field {
|
||||||
|
case WatchState:
|
||||||
|
val = ag.State()
|
||||||
|
default:
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
if val != current {
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
return val, true
|
||||||
|
}
|
||||||
|
if ctx.Err() != nil {
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
ag.changeCond.Wait()
|
||||||
|
if ctx.Err() != nil {
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// InitCond initializes the changeCond. Must be called once after construction.
|
||||||
|
func (ag *Agent) InitCond() {
|
||||||
|
ag.changeCond = sync.NewCond(&ag.changeMu)
|
||||||
|
}
|
||||||
|
|
||||||
|
// emit publishes an event on the agent's bus.
|
||||||
|
func (ag *Agent) emit(ev Event) {
|
||||||
|
ag.bus.Publish("event", ev)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CWD returns the agent's working directory.
|
||||||
|
func (ag *Agent) Cwd() string {
|
||||||
|
ag.stateMu.RLock()
|
||||||
|
c := ag.cwd
|
||||||
|
ag.stateMu.RUnlock()
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetCWD sets the agent's working directory (no validation — caller must validate).
|
||||||
|
func (ag *Agent) SetCwd(dir string) {
|
||||||
|
ag.stateMu.Lock()
|
||||||
|
ag.cwd = dir
|
||||||
|
ag.stateMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// effectiveCwd returns the agent's cwd, falling back to os.Getwd() if empty.
|
||||||
|
func (ag *Agent) effectiveCwd() string {
|
||||||
|
if c := ag.Cwd(); c != "" {
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
wd, _ := os.Getwd()
|
||||||
|
return wd
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsRunning returns true if the agent has an active turn in progress.
|
||||||
|
func (ag *Agent) IsRunning() bool {
|
||||||
|
return ag.currentAction.Load() != nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Interrupt cancels the current in-progress agent turn.
|
||||||
|
// Returns true if an action was running and was cancelled.
|
||||||
|
func (ag *Agent) Interrupt(cause error) bool {
|
||||||
|
if h := ag.currentAction.Load(); h != nil {
|
||||||
|
h.cancel(cause)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// Event is a typed output event emitted during an agent turn or in response
|
||||||
|
// to a command.
|
||||||
|
type Event struct {
|
||||||
|
Role string
|
||||||
|
Name string
|
||||||
|
Content string
|
||||||
|
ResponseID string
|
||||||
|
}
|
||||||
|
|
||||||
|
// EventHandler receives events from the agent.
|
||||||
|
type EventHandler func(Event)
|
||||||
|
|
||||||
|
// actionHandle holds the cancel function for the current agent turn.
|
||||||
|
type actionHandle struct {
|
||||||
|
cancel context.CancelCauseFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
// WatchField names supported by Agent.WaitChange.
|
||||||
|
const (
|
||||||
|
WatchState = "state"
|
||||||
|
)
|
||||||
|
|
||||||
|
// HasHistory returns true if the agent has an active conversation history.
|
||||||
|
func (ag *Agent) HasHistory() bool {
|
||||||
|
return ag.history != nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveTo saves the current history to the given path.
|
||||||
|
func (ag *Agent) SaveTo(path, name, cwd string) error {
|
||||||
|
if ag.history == nil {
|
||||||
|
return fmt.Errorf("no active session")
|
||||||
|
}
|
||||||
|
return ag.history.saveTo(path, name, ag.agentName, cwd)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Restore restores agent history from a persisted session.
|
||||||
|
func (ag *Agent) Restore(ps *PersistedAgent) {
|
||||||
|
ag.history = RestoreHistory(ps)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToolCallCount returns the total number of tool calls executed.
|
||||||
|
func (ag *Agent) ToolCallCount() int64 {
|
||||||
|
return ag.toolCallCount.Load()
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetSessionEnv injects session env vars into the execute server.
|
||||||
|
func (ag *Agent) SetSessionEnv(sessionID string) {
|
||||||
|
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
|
||||||
|
if es, ok := srv.(tools.EnvSetter); ok {
|
||||||
|
es.SetEnv("OLLIE_SESSION_ID", sessionID)
|
||||||
|
if ag.id != "" {
|
||||||
|
es.SetEnv("OLLIE_UNAME", ag.id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetEnv stores an environment variable on the execute server.
|
||||||
|
func (ag *Agent) SetEnv(key, value string) {
|
||||||
|
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
|
||||||
|
if es, ok := srv.(tools.EnvSetter); ok {
|
||||||
|
es.SetEnv(key, value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close releases agent resources (dispatcher, execute server).
|
||||||
|
func (ag *Agent) Close() {
|
||||||
|
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
|
||||||
|
if c, ok := srv.(interface{ Close() }); ok {
|
||||||
|
c.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetCWD updates the agent's working directory, preamble references, and dispatcher.
|
||||||
|
func (ag *Agent) SetCWD(dir string) {
|
||||||
|
oldCwd := ag.Cwd()
|
||||||
|
ag.SetCwd(dir)
|
||||||
|
if oldCwd != "" && dir != "" && oldCwd != dir {
|
||||||
|
ag.runtime.Preamble = strings.ReplaceAll(ag.runtime.Preamble, oldCwd, dir)
|
||||||
|
}
|
||||||
|
if ag.runtime != nil && ag.runtime.Dispatcher != nil {
|
||||||
|
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
|
||||||
|
if ws, ok := srv.(tools.CWDSetter); ok {
|
||||||
|
ws.SetCWD(dir)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
ag.notifyChange()
|
||||||
|
}
|
||||||
|
|
||||||
|
// RenamePreamble replaces old references in the preamble with new ones.
|
||||||
|
func (ag *Agent) RenamePreamble(old, new string) {
|
||||||
|
if ag.runtime != nil {
|
||||||
|
ag.runtime.Preamble = strings.ReplaceAll(ag.runtime.Preamble, old, new)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SaveFull persists the full session state (history + metadata) to the given path.
|
||||||
|
func (ag *Agent) SaveFull(path, sessionID, cwd, remote string) error {
|
||||||
|
if ag.history == nil {
|
||||||
|
return fmt.Errorf("no active session")
|
||||||
|
}
|
||||||
|
return ag.history.saveToFull(path, sessionID, ag.agentName,
|
||||||
|
ag.runtime.Backend.Name(), ag.runtime.Backend.Model(), cwd, remote)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CtxSz returns a human-readable context size string.
|
||||||
|
func (ag *Agent) CtxSz() string {
|
||||||
|
if ag.history == nil {
|
||||||
|
return "no active session"
|
||||||
|
}
|
||||||
|
ctxLen := ag.runtime.Backend.ContextLength(context.Background())
|
||||||
|
if ctxLen <= 0 {
|
||||||
|
ctxLen = defaultContextLength
|
||||||
|
}
|
||||||
|
estimated := ag.history.estimateTokens()
|
||||||
|
pct := estimated * 100 / ctxLen
|
||||||
|
return fmt.Sprintf("%d / %d (%d%%)", estimated, ctxLen, pct)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CostStr returns formatted cost information.
|
||||||
|
func (ag *Agent) CostStr() string {
|
||||||
|
if ag.history == nil {
|
||||||
|
return "no active session"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("costLast=$%.4f\ncostSession=$%.4f\n",
|
||||||
|
ag.history.LastTurnCostUSD, ag.history.SessionCostUSD)
|
||||||
|
}
|
||||||
|
|
||||||
|
// UsageStr returns formatted usage information.
|
||||||
|
func (ag *Agent) UsageStr() string {
|
||||||
|
if ag.history == nil {
|
||||||
|
return "no active session"
|
||||||
|
}
|
||||||
|
str := fmt.Sprintf("%d in, %d out, %d requests",
|
||||||
|
ag.history.TotalInputTokens, ag.history.TotalOutputTokens,
|
||||||
|
ag.history.TotalRequests)
|
||||||
|
if ag.history.TotalCachedInputTokens > 0 {
|
||||||
|
str += fmt.Sprintf(", %d cached", ag.history.TotalCachedInputTokens)
|
||||||
|
}
|
||||||
|
if ag.history.Estimated {
|
||||||
|
str += " [estimated]"
|
||||||
|
}
|
||||||
|
return str
|
||||||
|
}
|
||||||
|
|
||||||
|
// Context returns the full message context (system prompt + history).
|
||||||
|
func (ag *Agent) Context() []backend.Message {
|
||||||
|
var msgs []backend.Message
|
||||||
|
if ag.history != nil {
|
||||||
|
msgs = slices.Clone(ag.history.history())
|
||||||
|
}
|
||||||
|
if ag.runtime.Preamble != "" {
|
||||||
|
msgs = append([]backend.Message{{Role: "system", Content: ag.runtime.Preamble}}, msgs...)
|
||||||
|
}
|
||||||
|
return msgs
|
||||||
|
}
|
||||||
|
|
||||||
|
// SystemPrompt returns the rendered system prompt.
|
||||||
|
func (ag *Agent) SystemPrompt() string {
|
||||||
|
return ag.runtime.Preamble
|
||||||
|
}
|
||||||
|
|
||||||
|
// GenParams returns the current generation parameters.
|
||||||
|
func (ag *Agent) GenParams() backend.GenerationParams {
|
||||||
|
return ag.runtime.GenParams
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetGenParams sets the generation parameters.
|
||||||
|
func (ag *Agent) SetGenParams(params backend.GenerationParams) {
|
||||||
|
ag.runtime.GenParams = params
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompactionModel returns the configured compaction model name.
|
||||||
|
func (ag *Agent) CompactionModel() string {
|
||||||
|
return ag.runtime.CompactionModel
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetCompactionModel sets the compaction model.
|
||||||
|
func (ag *Agent) SetCompactionModel(model string) {
|
||||||
|
ag.runtime.CompactionModel = model
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListModels returns available models from the backend.
|
||||||
|
func (ag *Agent) ListModels() []string {
|
||||||
|
return ag.runtime.Backend.Models(context.Background())
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reactions returns a map of response ID → emoji for all recorded reactions.
|
||||||
|
func (ag *Agent) Reactions() map[string]string {
|
||||||
|
result := make(map[string]string)
|
||||||
|
if ag.history == nil {
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
for _, reaction := range ag.history.Reactions {
|
||||||
|
result[reaction.ResponseID] = reaction.Emoji
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// React records a reaction emoji against the most recent (or specified) response.
|
||||||
|
func (ag *Agent) React(responseID, emoji string) error {
|
||||||
|
if ag.history == nil {
|
||||||
|
return fmt.Errorf("no active session")
|
||||||
|
}
|
||||||
|
category, _, _ := classifyReaction(emoji)
|
||||||
|
if category == "unknown" {
|
||||||
|
return fmt.Errorf("unsupported reaction: %s", emoji)
|
||||||
|
}
|
||||||
|
if responseID == "" {
|
||||||
|
for i := len(ag.history.messages) - 1; i >= 0; i-- {
|
||||||
|
if ag.history.messages[i].Role == "assistant" {
|
||||||
|
responseID = ag.history.messages[i].ID
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if responseID == "" {
|
||||||
|
return fmt.Errorf("no assistant response to react to")
|
||||||
|
}
|
||||||
|
found := false
|
||||||
|
for i := range ag.history.messages {
|
||||||
|
if ag.history.messages[i].Role == "assistant" && ag.history.messages[i].ID == responseID {
|
||||||
|
found = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
return fmt.Errorf("assistant response not found: %s", responseID)
|
||||||
|
}
|
||||||
|
reaction := Reaction{ID: NewReactionID(), ResponseID: responseID, Emoji: emoji, Category: category, CreatedAt: time.Now()}
|
||||||
|
replaced := false
|
||||||
|
for i := range ag.history.Reactions {
|
||||||
|
if ag.history.Reactions[i].ResponseID == responseID {
|
||||||
|
if ag.history.Reactions[i].Emoji == emoji {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
ag.history.Reactions[i] = reaction
|
||||||
|
replaced = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !replaced {
|
||||||
|
ag.history.Reactions = append(ag.history.Reactions, reaction)
|
||||||
|
}
|
||||||
|
ag.history.recomputeReactionCounts()
|
||||||
|
ag.saveSession()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExecServer returns the execute server interface, or nil if unavailable.
|
||||||
|
// NOTE: This is temporary — detach operations should be proper Agent methods.
|
||||||
|
func (ag *Agent) ExecServer() interface{} {
|
||||||
|
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
srv, _ := ag.runtime.Dispatcher.GetServer("execute")
|
||||||
|
return srv
|
||||||
|
}
|
||||||
|
|
||||||
|
// Queue pushes a prompt onto the agent's FIFO.
|
||||||
|
func (ag *Agent) Queue(prompt string) {
|
||||||
|
ag.fifo.Push(prompt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// PopQueue pops the next prompt from the FIFO.
|
||||||
|
func (ag *Agent) PopQueue() (string, bool) {
|
||||||
|
return ag.fifo.Pop()
|
||||||
|
}
|
||||||
|
|
||||||
|
// BroadcastChange wakes all goroutines waiting on state changes.
|
||||||
|
func (ag *Agent) BroadcastChange() {
|
||||||
|
ag.changeMu.Lock()
|
||||||
|
ag.changeCond.Broadcast()
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// WaitForChange blocks until a state change is broadcast or ctx is cancelled.
|
||||||
|
func (ag *Agent) WaitForChange(ctx context.Context) {
|
||||||
|
ag.changeMu.Lock()
|
||||||
|
if ctx.Err() == nil {
|
||||||
|
ag.changeCond.Wait()
|
||||||
|
}
|
||||||
|
ag.changeMu.Unlock()
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
|
@ -0,0 +1,221 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"ollie/backend"
|
||||||
|
"ollie/tools"
|
||||||
|
)
|
||||||
|
|
||||||
|
// BuildRuntime constructs a Runtime from a pre-configured Dispatcher and
|
||||||
|
// optional agent config. cwd sets the working directory reported in the
|
||||||
|
// system prompt; if empty, the process working directory is used.
|
||||||
|
// env provides additional environment variables injected into prompt resolution
|
||||||
|
// subprocesses (e.g. OLLIE_SESSION_ID=xxx).
|
||||||
|
// The caller is responsible for registering all servers on d before calling this.
|
||||||
|
func BuildRuntime(cfg *AgentConfig, d tools.Dispatcher, cwd string, env []string, baseLayers ...string) *Runtime {
|
||||||
|
var messages []string
|
||||||
|
|
||||||
|
var allToolInfos []tools.ToolInfo
|
||||||
|
var allTools []backend.Tool
|
||||||
|
|
||||||
|
if cfg == nil || cfg.ToolsEnabled() {
|
||||||
|
var listErr error
|
||||||
|
allToolInfos, listErr = d.ListTools()
|
||||||
|
if listErr != nil {
|
||||||
|
messages = append(messages, fmt.Sprintf("list tools: %v", listErr))
|
||||||
|
}
|
||||||
|
// Only built-in executors (with InputSchema) become backend tools.
|
||||||
|
allTools = toolInfosToBackend(allToolInfos)
|
||||||
|
|
||||||
|
// Append named tool scripts for preamble listing only.
|
||||||
|
allToolInfos = append(allToolInfos, tools.DiscoverTools()...)
|
||||||
|
}
|
||||||
|
|
||||||
|
hooks := Hooks{}
|
||||||
|
var preamble string
|
||||||
|
var genParams backend.GenerationParams
|
||||||
|
var maxSteps int
|
||||||
|
if cfg != nil {
|
||||||
|
for k, v := range cfg.Hooks {
|
||||||
|
hooks[k] = []string(v)
|
||||||
|
}
|
||||||
|
if resolved, err := resolvePrompt(cfg.Prompt, cwd, env); err != nil {
|
||||||
|
fmt.Fprintf(os.Stderr, "resolve prompt: %v\n", err)
|
||||||
|
} else {
|
||||||
|
preamble = resolved
|
||||||
|
}
|
||||||
|
genParams = backend.GenerationParams{
|
||||||
|
MaxTokens: cfg.MaxTokens,
|
||||||
|
MaxCompletionTokens: cfg.MaxCompletionTokens,
|
||||||
|
Temperature: cfg.Temperature,
|
||||||
|
TopP: cfg.TopP,
|
||||||
|
TopK: cfg.TopK,
|
||||||
|
MinP: cfg.MinP,
|
||||||
|
TopA: cfg.TopA,
|
||||||
|
FrequencyPenalty: cfg.FrequencyPenalty,
|
||||||
|
PresencePenalty: cfg.PresencePenalty,
|
||||||
|
RepetitionPenalty: cfg.RepetitionPenalty,
|
||||||
|
ThinkingBudget: cfg.Reasoning,
|
||||||
|
ReasoningEffort: cfg.ReasoningEffort,
|
||||||
|
IncludeReasoning: cfg.IncludeReasoning,
|
||||||
|
ResponseFormat: cfg.ResponseFormat,
|
||||||
|
Stop: cfg.Stop,
|
||||||
|
Verbosity: cfg.Verbosity,
|
||||||
|
}
|
||||||
|
maxSteps = cfg.MaxSteps
|
||||||
|
if len(cfg.AllowTools) > 0 {
|
||||||
|
if srv, ok := d.GetServer("execute"); ok {
|
||||||
|
if rs, ok := srv.(tools.ToolRestrictionSetter); ok {
|
||||||
|
rs.SetAllowTools(cfg.AllowTools)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
exec := func(ctx context.Context, name string, args json.RawMessage) (string, []backend.ContentBlock, error) {
|
||||||
|
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)
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, err
|
||||||
|
}
|
||||||
|
text, blocks, isErr := extractToolResult(raw)
|
||||||
|
if isErr {
|
||||||
|
return "", nil, fmt.Errorf("%s", text)
|
||||||
|
}
|
||||||
|
return text, blocks, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var classify toolClassifier
|
||||||
|
if srv, ok := d.GetServer("execute"); ok {
|
||||||
|
if pc, ok := srv.(tools.ParallelClassifier); ok {
|
||||||
|
classify = pc.IsParallelRead
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var tierFn func(string, json.RawMessage) ResultTier
|
||||||
|
if srv, ok := d.GetServer("execute"); ok {
|
||||||
|
if tc, ok := srv.(tools.TierClassifier); ok {
|
||||||
|
tierFn = func(name string, args json.RawMessage) ResultTier {
|
||||||
|
switch tc.ResultTierArgs(name, args) {
|
||||||
|
case "cold":
|
||||||
|
return TierCold
|
||||||
|
case "warm":
|
||||||
|
return TierWarm
|
||||||
|
default:
|
||||||
|
return TierHot
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var backendName, modelName, compactionModel string
|
||||||
|
if cfg != nil {
|
||||||
|
backendName = cfg.Backend
|
||||||
|
modelName = cfg.Model
|
||||||
|
compactionModel = cfg.CompactionModel
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prepend base layers (system prompt, operational model, environment)
|
||||||
|
// before the agent-specific preamble.
|
||||||
|
if len(baseLayers) > 0 {
|
||||||
|
var prefix strings.Builder
|
||||||
|
for _, layer := range baseLayers {
|
||||||
|
if layer != "" {
|
||||||
|
prefix.WriteString(layer)
|
||||||
|
prefix.WriteByte('\n')
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if prefix.Len() > 0 {
|
||||||
|
preamble = prefix.String() + preamble
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Append compact tool surface listing (name + description).
|
||||||
|
var toolListing strings.Builder
|
||||||
|
for _, ti := range allToolInfos {
|
||||||
|
if ti.Description != "" && ti.Server == "" {
|
||||||
|
fmt.Fprintf(&toolListing, "- **%s** — %s\n", ti.Name, ti.Description)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if toolListing.Len() > 0 {
|
||||||
|
preamble += "\n# Available Tools\n\n" + toolListing.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Runtime{
|
||||||
|
Dispatcher: d,
|
||||||
|
Tools: allTools,
|
||||||
|
Exec: exec,
|
||||||
|
ClassifyTool: classify,
|
||||||
|
ClassifyTier: tierFn,
|
||||||
|
Hooks: hooks,
|
||||||
|
Preamble: preamble,
|
||||||
|
GenParams: genParams,
|
||||||
|
MaxSteps: maxSteps,
|
||||||
|
CfgBackend: backendName,
|
||||||
|
CfgModel: modelName,
|
||||||
|
CompactionModel: compactionModel,
|
||||||
|
Messages: messages,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func toolInfosToBackend(infos []tools.ToolInfo) []backend.Tool {
|
||||||
|
out := make([]backend.Tool, len(infos))
|
||||||
|
for i, t := range infos {
|
||||||
|
out[i] = backend.Tool{
|
||||||
|
Name: t.Name,
|
||||||
|
Description: t.Description,
|
||||||
|
Parameters: t.InputSchema,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func extractToolResult(raw json.RawMessage) (text string, contentBlocks []backend.ContentBlock, isError bool) {
|
||||||
|
var result struct {
|
||||||
|
IsError bool `json:"isError"`
|
||||||
|
Content []struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Text string `json:"text"`
|
||||||
|
MediaType string `json:"media_type"`
|
||||||
|
Data string `json:"data"`
|
||||||
|
} `json:"content"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(raw, &result); err != nil {
|
||||||
|
return string(raw), nil, false
|
||||||
|
}
|
||||||
|
var parts []string
|
||||||
|
for _, c := range result.Content {
|
||||||
|
switch c.Type {
|
||||||
|
case "text":
|
||||||
|
parts = append(parts, c.Text)
|
||||||
|
case "image":
|
||||||
|
contentBlocks = append(contentBlocks, backend.ContentBlock{
|
||||||
|
Type: "image",
|
||||||
|
ImageSource: &backend.ImageSource{
|
||||||
|
Type: "base64",
|
||||||
|
MediaType: c.MediaType,
|
||||||
|
Data: c.Data,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return strings.Join(parts, "\n"), contentBlocks, result.IsError
|
||||||
|
}
|
||||||
|
|
@ -0,0 +1,364 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// HandleCommand processes agent-level slash commands. Returns true if the
|
||||||
|
// input was recognized as a command, false otherwise.
|
||||||
|
// The emit function is used to send output events. sessionID and agentsDir
|
||||||
|
// are passed from the session to avoid the agent needing session-level knowledge.
|
||||||
|
func (ag *Agent) HandleCommand(ctx context.Context, input string) bool {
|
||||||
|
if !strings.HasPrefix(input, "/") {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
parts := strings.Fields(input)
|
||||||
|
if len(parts) == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
cmd := parts[0]
|
||||||
|
args := parts[1:]
|
||||||
|
|
||||||
|
switch cmd {
|
||||||
|
case "/i":
|
||||||
|
ag.cmdInject(ctx, args)
|
||||||
|
case "/irw":
|
||||||
|
ag.cmdInjectRewrite(ctx, args)
|
||||||
|
case "/backend":
|
||||||
|
ag.cmdBackend(ctx, args)
|
||||||
|
case "/models":
|
||||||
|
ag.cmdModels(ctx, args)
|
||||||
|
case "/model":
|
||||||
|
ag.cmdModel(ctx, args)
|
||||||
|
case "/maxsteps":
|
||||||
|
ag.cmdMaxSteps(ctx, args)
|
||||||
|
case "/agents":
|
||||||
|
ag.cmdAgents(ctx, args)
|
||||||
|
case "/agent":
|
||||||
|
ag.cmdAgent(ctx, args)
|
||||||
|
case "/compact":
|
||||||
|
ag.cmdCompact(ctx, args)
|
||||||
|
case "/context":
|
||||||
|
ag.cmdContext(ctx, args)
|
||||||
|
case "/cost":
|
||||||
|
ag.cmdCost(ctx, args)
|
||||||
|
case "/usage":
|
||||||
|
ag.cmdUsage(ctx, args)
|
||||||
|
case "/history":
|
||||||
|
ag.cmdHistory(ctx, args)
|
||||||
|
case "/clear":
|
||||||
|
ag.cmdClear(ctx, args)
|
||||||
|
case "/sp":
|
||||||
|
ag.cmdSP(ctx, args)
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdBackend(_ context.Context, args []string) {
|
||||||
|
if len(args) == 0 {
|
||||||
|
ag.emit(infoEvent(ag.runtime.Backend.Name()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if ag.IsRunning() {
|
||||||
|
ag.emit(infoEvent("error: cannot switch backend while agent is running"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
be, err := ag.newBackend(args[0])
|
||||||
|
if err != nil {
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("error: failed to switch backend: %v", err)))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.runtime.Backend = be
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("switched backend to: %s (model: %s)", be.Name(), be.Model())))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdModels(ctx context.Context, args []string) {
|
||||||
|
models := ag.runtime.Backend.Models(ctx)
|
||||||
|
if len(models) == 0 {
|
||||||
|
ag.emit(infoEvent("no models available"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
slices.Sort(models)
|
||||||
|
current := ag.runtime.Backend.Model()
|
||||||
|
for _, m := range models {
|
||||||
|
marker := " "
|
||||||
|
if m == current {
|
||||||
|
marker = "* "
|
||||||
|
}
|
||||||
|
ag.emit(infoEvent(marker + m))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdModel(_ context.Context, args []string) {
|
||||||
|
if len(args) == 0 {
|
||||||
|
ag.emit(infoEvent(ag.runtime.Backend.Model()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.runtime.Backend.SetModel(args[0])
|
||||||
|
ag.emit(infoEvent("switched model to: " + args[0]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdMaxSteps(_ context.Context, args []string) {
|
||||||
|
if len(args) == 0 {
|
||||||
|
if ag.runtime.MaxSteps == 0 {
|
||||||
|
ag.emit(infoEvent("maxsteps: unlimited"))
|
||||||
|
} else {
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("maxsteps: %d", ag.runtime.MaxSteps)))
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
n, err := strconv.Atoi(args[0])
|
||||||
|
if err != nil || n < 0 {
|
||||||
|
ag.emit(infoEvent("error: maxsteps requires a non-negative integer (0 = unlimited)"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.runtime.MaxSteps = n
|
||||||
|
if n == 0 {
|
||||||
|
ag.emit(infoEvent("maxsteps: unlimited"))
|
||||||
|
} else {
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("maxsteps: %d", n)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdAgents(_ context.Context, _ []string) {
|
||||||
|
seen := make(map[string]bool)
|
||||||
|
found := false
|
||||||
|
for _, dir := range AgentsDirs() {
|
||||||
|
entries, err := os.ReadDir(dir)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, e := range entries {
|
||||||
|
if e.IsDir() || !strings.HasSuffix(e.Name(), ".json") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
name := strings.TrimSuffix(e.Name(), ".json")
|
||||||
|
if seen[name] {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[name] = true
|
||||||
|
marker := " "
|
||||||
|
if name == ag.agentName {
|
||||||
|
marker = "* "
|
||||||
|
}
|
||||||
|
ag.emit(infoEvent(marker + name))
|
||||||
|
found = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !found {
|
||||||
|
ag.emit(infoEvent("no agents found"))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdAgent(_ context.Context, args []string) {
|
||||||
|
if len(args) == 0 {
|
||||||
|
ag.emit(infoEvent("active agent: " + ag.agentName))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if ag.IsRunning() {
|
||||||
|
ag.emit(infoEvent("error: cannot switch agent while agent is running"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
name := args[0]
|
||||||
|
cfgPath := AgentConfigPath(ag.agentsDir, name)
|
||||||
|
f, err := os.Open(cfgPath)
|
||||||
|
if err != nil {
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err)))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cfg, err := Load(f)
|
||||||
|
f.Close()
|
||||||
|
if err != nil {
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err)))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
d := ag.newDispatcher()
|
||||||
|
if d == nil {
|
||||||
|
ag.emit(infoEvent("error: no dispatcher configured"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
env := []string{"OLLIE_SESSION_ID=" + ag.sessionID, "OLLIE_UNAME=" + ag.id}
|
||||||
|
env = append(env, ag.promptEnvExtra...)
|
||||||
|
rt := BuildRuntime(cfg, d, ag.cwd, env, ag.baseLayers...)
|
||||||
|
if rt.CfgBackend != "" {
|
||||||
|
newBe, err := ag.newBackend(rt.CfgBackend)
|
||||||
|
if err != nil {
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("error: backend %q: %v", rt.CfgBackend, err)))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if rt.CfgModel != "" {
|
||||||
|
newBe.SetModel(rt.CfgModel)
|
||||||
|
}
|
||||||
|
rt.Backend = newBe
|
||||||
|
} else {
|
||||||
|
rt.Backend = ag.runtime.Backend
|
||||||
|
if rt.CfgModel != "" {
|
||||||
|
rt.Backend.SetModel(rt.CfgModel)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
ag.runtime = rt
|
||||||
|
ag.agentName = name
|
||||||
|
ag.history = nil
|
||||||
|
ag.notifyChange()
|
||||||
|
for _, msg := range rt.Messages {
|
||||||
|
ag.emit(infoEvent(msg))
|
||||||
|
}
|
||||||
|
ag.emit(infoEvent("agent: " + name))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdCompact(ctx context.Context, _ []string) {
|
||||||
|
if ag.IsRunning() {
|
||||||
|
ag.emit(infoEvent("error: cannot compact while agent is running"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if ag.history == nil {
|
||||||
|
ag.emit(infoEvent("nothing to compact"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.SetState("compacting")
|
||||||
|
n, err := ag.runCompact(ctx, "manual")
|
||||||
|
ag.SetState("idle")
|
||||||
|
if err != nil {
|
||||||
|
ag.emit(infoEvent("compact error: " + err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if n == 0 {
|
||||||
|
ag.emit(infoEvent("nothing to compact"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("compacted %d messages", n)))
|
||||||
|
ag.saveSession()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdContext(ctx context.Context, _ []string) {
|
||||||
|
if ag.history == nil {
|
||||||
|
ag.emit(infoEvent("no active session"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ctxLen := ag.runtime.Backend.ContextLength(ctx)
|
||||||
|
if ctxLen <= 0 {
|
||||||
|
ctxLen = defaultContextLength
|
||||||
|
}
|
||||||
|
estimated := ag.history.estimateTokens()
|
||||||
|
pct := estimated * 100 / ctxLen
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("~%d / %d tokens (%d%%)", estimated, ctxLen, pct)))
|
||||||
|
ag.emit(infoEvent(strings.TrimRight(ag.history.contextDebug(), "\n")))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdCost(_ context.Context, _ []string) {
|
||||||
|
if ag.history == nil {
|
||||||
|
ag.emit(infoEvent("no active session"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("last=$%.4f session=$%.4f",
|
||||||
|
ag.history.LastTurnCostUSD, ag.history.SessionCostUSD)))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdUsage(ctx context.Context, _ []string) {
|
||||||
|
if ag.history == nil {
|
||||||
|
ag.emit(infoEvent("no active session"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ctxLen := ag.runtime.Backend.ContextLength(ctx)
|
||||||
|
if ctxLen <= 0 {
|
||||||
|
ctxLen = defaultContextLength
|
||||||
|
}
|
||||||
|
estimated := ag.history.estimateTokens()
|
||||||
|
pct := estimated * 100 / ctxLen
|
||||||
|
usageStr := fmt.Sprintf("~%d / %d tokens (%d%%) | %d in, %d out, %d requests",
|
||||||
|
estimated, ctxLen, pct,
|
||||||
|
ag.history.TotalInputTokens, ag.history.TotalOutputTokens,
|
||||||
|
ag.history.TotalRequests)
|
||||||
|
if ag.history.Estimated {
|
||||||
|
usageStr += " [estimated]"
|
||||||
|
}
|
||||||
|
ag.emit(infoEvent(usageStr))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdHistory(_ context.Context, _ []string) {
|
||||||
|
if ag.history == nil {
|
||||||
|
ag.emit(infoEvent("no active session"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, msg := range ag.history.history() {
|
||||||
|
preview := msg.Content
|
||||||
|
if len(preview) > 200 {
|
||||||
|
preview = preview[:200] + "..."
|
||||||
|
}
|
||||||
|
ag.emit(infoEvent(fmt.Sprintf("[%s] %s", msg.Role, preview)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdClear(_ context.Context, _ []string) {
|
||||||
|
if ag.IsRunning() {
|
||||||
|
ag.emit(infoEvent("error: cannot clear while agent is running"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.history = nil
|
||||||
|
ag.emit(infoEvent("cleared"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdSP(_ context.Context, _ []string) {
|
||||||
|
ag.emit(infoEvent(ag.runtime.Preamble))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdInject(_ context.Context, args []string) {
|
||||||
|
prompt := strings.Join(args, " ")
|
||||||
|
if prompt == "" {
|
||||||
|
ag.emit(infoEvent("error: /i requires a prompt"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if ag.IsRunning() {
|
||||||
|
ag.Inject(prompt)
|
||||||
|
} else {
|
||||||
|
go ag.Submit(context.Background(), prompt)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (ag *Agent) cmdInjectRewrite(_ context.Context, args []string) {
|
||||||
|
prompt := strings.Join(args, " ")
|
||||||
|
if prompt == "" {
|
||||||
|
ag.emit(infoEvent("error: /irw requires a prompt"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.InjectRewrite(prompt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Inject queues a prompt for mid-turn injection.
|
||||||
|
func (ag *Agent) Inject(prompt string) {
|
||||||
|
if !ag.pendingInject.CompareAndSwap(nil, &prompt) {
|
||||||
|
ag.fifo.Push(prompt)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ag.emit(Event{Role: "info", Content: "\n"})
|
||||||
|
ag.emit(Event{Role: "user", Content: prompt})
|
||||||
|
}
|
||||||
|
|
||||||
|
// InjectRewrite replaces the pending inject.
|
||||||
|
func (ag *Agent) InjectRewrite(prompt string) {
|
||||||
|
ag.pendingInject.Store(&prompt)
|
||||||
|
ag.emit(Event{Role: "info", Content: "\n"})
|
||||||
|
ag.emit(Event{Role: "user", Content: prompt})
|
||||||
|
}
|
||||||
|
|
||||||
|
// CompactionSnapshot returns the pre-compaction snapshot for external persistence.
|
||||||
|
// Returns nil if no history exists.
|
||||||
|
func (ag *Agent) CompactionSnapshot() json.RawMessage {
|
||||||
|
if ag.history == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
snap := ag.history.PreCompactionSnapshot()
|
||||||
|
data, err := json.Marshal(snap)
|
||||||
|
if err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return data
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import "os"
|
import "os"
|
||||||
|
|
||||||
|
|
@ -0,0 +1,41 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"ollie/paths"
|
||||||
|
)
|
||||||
|
|
||||||
|
// DefaultPromptsDir returns the default directory for prompt templates.
|
||||||
|
func DefaultPromptsDir() string {
|
||||||
|
return paths.CfgDir() + "/prompts"
|
||||||
|
}
|
||||||
|
|
||||||
|
// PromptsDirs returns all prompt directories from OLLIE_PROMPTS_PATH (colon-separated).
|
||||||
|
func PromptsDirs() []string {
|
||||||
|
if p := os.Getenv("OLLIE_PROMPTS_PATH"); p != "" {
|
||||||
|
return strings.Split(p, ":")
|
||||||
|
}
|
||||||
|
return []string{DefaultPromptsDir()}
|
||||||
|
}
|
||||||
|
|
||||||
|
// AgentsDirs returns all agent directories from OLLIE_AGENTS_PATH (colon-separated).
|
||||||
|
func AgentsDirs() []string {
|
||||||
|
if p := os.Getenv("OLLIE_AGENTS_PATH"); p != "" {
|
||||||
|
return strings.Split(p, ":")
|
||||||
|
}
|
||||||
|
return []string{paths.CfgDir() + "/agents"}
|
||||||
|
}
|
||||||
|
|
||||||
|
// AgentConfigPath resolves the config file path for a named agent.
|
||||||
|
// It searches all directories in OLLIE_AGENTS_PATH, falling back to agentsDir.
|
||||||
|
func AgentConfigPath(agentsDir, name string) string {
|
||||||
|
for _, dir := range AgentsDirs() {
|
||||||
|
p := dir + "/" + name + ".json"
|
||||||
|
if _, err := os.Stat(p); err == nil {
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return agentsDir + "/" + name + ".json"
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"strings"
|
"strings"
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import "sync"
|
import "sync"
|
||||||
|
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"sync"
|
"sync"
|
||||||
|
|
@ -1,13 +1,15 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"crypto/rand"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
"strconv"
|
||||||
|
|
||||||
"ollie/backend"
|
"ollie/backend"
|
||||||
)
|
)
|
||||||
|
|
@ -694,3 +696,35 @@ func (s *History) contextDebug() string {
|
||||||
}
|
}
|
||||||
return sb.String()
|
return sb.String()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// NewResponseID generates a unique identifier for a single assistant response.
|
||||||
|
func NewResponseID() string {
|
||||||
|
b := make([]byte, 3)
|
||||||
|
rand.Read(b) //nolint:errcheck
|
||||||
|
return "resp_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewReactionID generates a unique identifier for a user reaction.
|
||||||
|
func NewReactionID() string {
|
||||||
|
b := make([]byte, 3)
|
||||||
|
rand.Read(b) //nolint:errcheck
|
||||||
|
return "react_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// classifyReaction returns a category and description for a reaction emoji.
|
||||||
|
func classifyReaction(emoji string) (category, description string, positive bool) {
|
||||||
|
switch emoji {
|
||||||
|
case "👍", "✅":
|
||||||
|
return "positive", "The response was good. Keep doing what you're doing.", true
|
||||||
|
case "🚀", "🎉":
|
||||||
|
return "excellent", "The response was exactly what was wanted.", true
|
||||||
|
case "👎", "❌":
|
||||||
|
return "negative", "The response was wrong or unhelpful.", false
|
||||||
|
case "💩", "🤬":
|
||||||
|
return "terrible", "The response was fundamentally wrong. Stop this approach entirely and reassess from scratch.", false
|
||||||
|
case "🤔":
|
||||||
|
return "confused", "The response was unclear or confusing.", false
|
||||||
|
default:
|
||||||
|
return "unknown", "", false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
@ -1037,3 +1037,10 @@ func retryCountdown(ctx context.Context, cfg agentConfig, wait time.Duration) er
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// defaultToolResultMaxBytes caps tool result content sent back to the model.
|
||||||
|
const defaultToolResultMaxBytes = 131072
|
||||||
|
|
||||||
|
// toolClassifier reports whether a named tool is safe to run concurrently
|
||||||
|
// with other read-class tools. nil means treat all tools as serial.
|
||||||
|
type toolClassifier func(name string) bool
|
||||||
|
|
@ -0,0 +1,60 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/simonfxr/pubsub"
|
||||||
|
"ollie/backend"
|
||||||
|
olog "ollie/log"
|
||||||
|
"ollie/tools"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AgentCfg is the configuration for constructing a new Agent.
|
||||||
|
type AgentCfg struct {
|
||||||
|
History *History
|
||||||
|
Runtime *Runtime
|
||||||
|
AgentName string
|
||||||
|
AgentsDir string
|
||||||
|
AgentID string // unique agent identity
|
||||||
|
CWD string // working directory for tool execution
|
||||||
|
BaseLayers []string
|
||||||
|
PromptEnvExtra []string
|
||||||
|
NewDispatcher func() tools.Dispatcher
|
||||||
|
NewBackend func(string) (backend.Backend, error)
|
||||||
|
Bus *pubsub.Bus
|
||||||
|
Log *olog.Logger
|
||||||
|
AuditLog *olog.Logger
|
||||||
|
SessionID string
|
||||||
|
StartupMsgs []string
|
||||||
|
ReadPlanStep func() string
|
||||||
|
SaveSession func()
|
||||||
|
FlushSave func()
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAgent constructs an Agent from the given configuration.
|
||||||
|
func NewAgent(cfg AgentCfg) *Agent {
|
||||||
|
ag := &Agent{
|
||||||
|
history: cfg.History,
|
||||||
|
runtime: cfg.Runtime,
|
||||||
|
agentName: cfg.AgentName,
|
||||||
|
agentsDir: cfg.AgentsDir,
|
||||||
|
id: cfg.AgentID,
|
||||||
|
cwd: cfg.CWD,
|
||||||
|
baseLayers: cfg.BaseLayers,
|
||||||
|
promptEnvExtra: cfg.PromptEnvExtra,
|
||||||
|
newDispatcher: cfg.NewDispatcher,
|
||||||
|
newBackend: cfg.NewBackend,
|
||||||
|
bus: cfg.Bus,
|
||||||
|
log: cfg.Log,
|
||||||
|
auditLog: cfg.AuditLog,
|
||||||
|
sessionID: cfg.SessionID,
|
||||||
|
startupMessages: cfg.StartupMsgs,
|
||||||
|
readPlanStep: cfg.ReadPlanStep,
|
||||||
|
saveSession: cfg.SaveSession,
|
||||||
|
flushSave: cfg.FlushSave,
|
||||||
|
state: "idle",
|
||||||
|
}
|
||||||
|
ag.changeCond = sync.NewCond(&ag.changeMu)
|
||||||
|
ag.turnError = ag.defaultTurnError
|
||||||
|
return ag
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
@ -464,3 +464,29 @@ func (ag *Agent) runCompact(ctx context.Context, trigger string) (int, error) {
|
||||||
}
|
}
|
||||||
return n, nil
|
return n, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ErrInterrupted is returned when the user cancels an agent turn (Ctrl-C).
|
||||||
|
var ErrInterrupted = errors.New("interrupted")
|
||||||
|
|
||||||
|
// defaultContextLength is used when the backend cannot report the model's
|
||||||
|
// actual context window. 128k tokens is a safe default for modern models.
|
||||||
|
const defaultContextLength = 128000
|
||||||
|
|
||||||
|
// infoEvent wraps a plain-text message as an info Event.
|
||||||
|
func infoEvent(text string) Event {
|
||||||
|
return Event{Role: "info", Content: text + "\n"}
|
||||||
|
}
|
||||||
|
|
||||||
|
// firstSentence returns the first sentence of s (up to the first period or
|
||||||
|
// newline), trimmed. Falls back to s truncated at 80 chars if no sentence end found.
|
||||||
|
func firstSentence(s string) string {
|
||||||
|
for i, r := range s {
|
||||||
|
if r == '.' || r == '\n' {
|
||||||
|
return strings.TrimSpace(s[:i+1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(s) > 80 {
|
||||||
|
return s[:77] + "..."
|
||||||
|
}
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
package session
|
package agent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
210
session/agent.go
210
session/agent.go
|
|
@ -1,210 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
|
|
||||||
"github.com/simonfxr/pubsub"
|
|
||||||
"ollie/backend"
|
|
||||||
olog "ollie/log"
|
|
||||||
"ollie/tools"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Agent holds the state of the current Agent entity (the "agent"
|
|
||||||
// in the traditional sense). It is swappable: when the user runs /agent,
|
|
||||||
// a new Agent is built from the new agent config while the session
|
|
||||||
// host remains stable.
|
|
||||||
type Agent struct {
|
|
||||||
history *History
|
|
||||||
runtime *Runtime
|
|
||||||
cfg agentConfig // per-turn config built from runtime
|
|
||||||
agentName string
|
|
||||||
agentsDir string
|
|
||||||
baseLayers []string // system prompt layers for /agent reloads
|
|
||||||
promptEnvExtra []string // PRIME_* vars for prompt resolution
|
|
||||||
newDispatcher func() tools.Dispatcher
|
|
||||||
newBackend func(string) (backend.Backend, error)
|
|
||||||
currentAction atomic.Pointer[actionHandle]
|
|
||||||
warnedContext bool
|
|
||||||
resultCache sync.Map
|
|
||||||
|
|
||||||
// Execution state — owned by the agent, protected by stateMu.
|
|
||||||
state string // "idle", "thinking", "calling: <tool>"
|
|
||||||
reply string // last assistant response
|
|
||||||
cwd string // working directory for tool execution
|
|
||||||
id string // agent identity (unique principal)
|
|
||||||
fifo Fifo // prompt queue
|
|
||||||
toolCallCount atomic.Int64
|
|
||||||
pendingInject atomic.Pointer[string]
|
|
||||||
submitMu sync.Mutex // serializes Submit calls (commands + turns)
|
|
||||||
stateMu sync.RWMutex
|
|
||||||
changeMu sync.Mutex
|
|
||||||
changeCond *sync.Cond
|
|
||||||
|
|
||||||
// Injected session-level dependencies (set at creation, stable for agent lifetime).
|
|
||||||
bus *pubsub.Bus
|
|
||||||
log *olog.Logger
|
|
||||||
auditLog *olog.Logger
|
|
||||||
sessionID string // the owning session's ID
|
|
||||||
startupMessages []string
|
|
||||||
readPlanStep func() string
|
|
||||||
saveSession func() // trigger debounced persistence
|
|
||||||
flushSave func() // immediately flush persistence
|
|
||||||
turnError func(ctx context.Context, errType, errMsg string) HookResult // overridable for tests
|
|
||||||
}
|
|
||||||
|
|
||||||
// Backend returns the active backend from the runtime.
|
|
||||||
func (ag *Agent) Backend() backend.Backend {
|
|
||||||
if ag.runtime == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return ag.runtime.Backend
|
|
||||||
}
|
|
||||||
|
|
||||||
// Name returns the agent's name.
|
|
||||||
func (ag *Agent) Name() string { return ag.agentName }
|
|
||||||
|
|
||||||
// ID returns the agent's unique identity.
|
|
||||||
func (ag *Agent) ID() string { return ag.id }
|
|
||||||
|
|
||||||
// BackendName returns the name of the active backend.
|
|
||||||
func (ag *Agent) BackendName() string {
|
|
||||||
if ag.runtime == nil || ag.runtime.Backend == nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
return ag.runtime.Backend.Name()
|
|
||||||
}
|
|
||||||
|
|
||||||
// ModelName returns the name of the active model.
|
|
||||||
func (ag *Agent) ModelName() string {
|
|
||||||
if ag.runtime == nil || ag.runtime.Backend == nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
return ag.runtime.Backend.Model()
|
|
||||||
}
|
|
||||||
|
|
||||||
// State returns the agent's current execution state.
|
|
||||||
func (ag *Agent) State() string {
|
|
||||||
ag.stateMu.RLock()
|
|
||||||
s := ag.state
|
|
||||||
ag.stateMu.RUnlock()
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetState sets the agent's execution state and notifies waiters.
|
|
||||||
func (ag *Agent) SetState(state string) {
|
|
||||||
ag.stateMu.Lock()
|
|
||||||
ag.state = state
|
|
||||||
ag.stateMu.Unlock()
|
|
||||||
ag.notifyChange()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Reply returns the agent's last assistant response.
|
|
||||||
func (ag *Agent) Reply() string {
|
|
||||||
ag.stateMu.RLock()
|
|
||||||
r := ag.reply
|
|
||||||
ag.stateMu.RUnlock()
|
|
||||||
return r
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetReply sets the agent's last response.
|
|
||||||
func (ag *Agent) SetReply(reply string) {
|
|
||||||
ag.stateMu.Lock()
|
|
||||||
ag.reply = reply
|
|
||||||
ag.stateMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
// notifyChange wakes all goroutines waiting on state changes.
|
|
||||||
func (ag *Agent) notifyChange() {
|
|
||||||
ag.changeMu.Lock()
|
|
||||||
ag.changeCond.Broadcast()
|
|
||||||
ag.changeMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
// WaitChange blocks until the agent's state differs from current.
|
|
||||||
// Returns the new value and true, or ("", false) if ctx is cancelled.
|
|
||||||
func (ag *Agent) WaitChange(ctx context.Context, field, current string) (string, bool) {
|
|
||||||
done := make(chan struct{})
|
|
||||||
context.AfterFunc(ctx, func() {
|
|
||||||
ag.changeMu.Lock()
|
|
||||||
ag.changeCond.Broadcast()
|
|
||||||
ag.changeMu.Unlock()
|
|
||||||
close(done)
|
|
||||||
})
|
|
||||||
|
|
||||||
ag.changeMu.Lock()
|
|
||||||
for {
|
|
||||||
var val string
|
|
||||||
switch field {
|
|
||||||
case WatchState:
|
|
||||||
val = ag.State()
|
|
||||||
default:
|
|
||||||
ag.changeMu.Unlock()
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
if val != current {
|
|
||||||
ag.changeMu.Unlock()
|
|
||||||
return val, true
|
|
||||||
}
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
ag.changeMu.Unlock()
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
ag.changeCond.Wait()
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
ag.changeMu.Unlock()
|
|
||||||
return "", false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// InitCond initializes the changeCond. Must be called once after construction.
|
|
||||||
func (ag *Agent) InitCond() {
|
|
||||||
ag.changeCond = sync.NewCond(&ag.changeMu)
|
|
||||||
}
|
|
||||||
|
|
||||||
// emit publishes an event on the agent's bus.
|
|
||||||
func (ag *Agent) emit(ev Event) {
|
|
||||||
ag.bus.Publish("event", ev)
|
|
||||||
}
|
|
||||||
|
|
||||||
// CWD returns the agent's working directory.
|
|
||||||
func (ag *Agent) Cwd() string {
|
|
||||||
ag.stateMu.RLock()
|
|
||||||
c := ag.cwd
|
|
||||||
ag.stateMu.RUnlock()
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
|
|
||||||
// SetCWD sets the agent's working directory (no validation — caller must validate).
|
|
||||||
func (ag *Agent) SetCwd(dir string) {
|
|
||||||
ag.stateMu.Lock()
|
|
||||||
ag.cwd = dir
|
|
||||||
ag.stateMu.Unlock()
|
|
||||||
}
|
|
||||||
|
|
||||||
// effectiveCwd returns the agent's cwd, falling back to os.Getwd() if empty.
|
|
||||||
func (ag *Agent) effectiveCwd() string {
|
|
||||||
if c := ag.Cwd(); c != "" {
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
wd, _ := os.Getwd()
|
|
||||||
return wd
|
|
||||||
}
|
|
||||||
|
|
||||||
// IsRunning returns true if the agent has an active turn in progress.
|
|
||||||
func (ag *Agent) IsRunning() bool {
|
|
||||||
return ag.currentAction.Load() != nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Interrupt cancels the current in-progress agent turn.
|
|
||||||
// Returns true if an action was running and was cancelled.
|
|
||||||
func (ag *Agent) Interrupt(cause error) bool {
|
|
||||||
if h := ag.currentAction.Load(); h != nil {
|
|
||||||
h.cancel(cause)
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
@ -6,10 +6,9 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"slices"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"ollie/agent"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (a *Session) handleCommand(ctx context.Context, input string) bool {
|
func (a *Session) handleCommand(ctx context.Context, input string) bool {
|
||||||
|
|
@ -24,292 +23,21 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
|
||||||
cmd := parts[0]
|
cmd := parts[0]
|
||||||
args := parts[1:]
|
args := parts[1:]
|
||||||
|
|
||||||
|
// Try agent-level commands first (they access agent internals directly).
|
||||||
|
if a.r.HandleCommand(ctx, input) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
listFromHandler := func(name string) {
|
listFromHandler := func(name string) {
|
||||||
if h := a.listHandlers[name]; h != nil {
|
if h := a.listHandlers[name]; h != nil {
|
||||||
for _, item := range h() {
|
for _, item := range h() {
|
||||||
a.emit(infoEvent(" " + item))
|
a.emit(agent.Event{Role: "info", Content: " " + item + "\n"})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
type cmdFn func([]string)
|
type cmdFn func([]string)
|
||||||
cmds := map[string]cmdFn{
|
cmds := map[string]cmdFn{
|
||||||
"/i": func(args []string) {
|
|
||||||
prompt := strings.Join(args, " ")
|
|
||||||
if prompt == "" {
|
|
||||||
a.emit(infoEvent("error: /i requires a prompt"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if a.IsRunning() {
|
|
||||||
a.Inject(prompt)
|
|
||||||
} else {
|
|
||||||
go a.Submit(context.Background(), prompt)
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
"/irw": func(args []string) {
|
|
||||||
prompt := strings.Join(args, " ")
|
|
||||||
if prompt == "" {
|
|
||||||
a.emit(infoEvent("error: /irw requires a prompt"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
a.injectRewrite(prompt)
|
|
||||||
},
|
|
||||||
|
|
||||||
"/backend": func(args []string) {
|
|
||||||
if len(args) == 0 {
|
|
||||||
a.emit(infoEvent(a.r.runtime.Backend.Name()))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if a.IsRunning() {
|
|
||||||
a.emit(infoEvent("error: cannot switch backend while agent is running"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
be, err := a.r.newBackend(args[0])
|
|
||||||
if err != nil {
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("error: failed to switch backend: %v", err)))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
a.r.runtime.Backend = be
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("switched backend to: %s (model: %s)", be.Name(), be.Model())))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/models": func(args []string) {
|
|
||||||
models := a.r.runtime.Backend.Models(ctx)
|
|
||||||
if len(models) == 0 {
|
|
||||||
a.emit(infoEvent("no models available"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
slices.Sort(models)
|
|
||||||
current := a.r.runtime.Backend.Model()
|
|
||||||
for _, m := range models {
|
|
||||||
marker := " "
|
|
||||||
if m == current {
|
|
||||||
marker = "* "
|
|
||||||
}
|
|
||||||
a.emit(infoEvent(marker + m))
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
"/model": func(args []string) {
|
|
||||||
if len(args) == 0 {
|
|
||||||
a.emit(infoEvent(a.r.runtime.Backend.Model()))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
a.r.runtime.Backend.SetModel(args[0])
|
|
||||||
a.emit(infoEvent("switched model to: " + args[0]))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/maxsteps": func(args []string) {
|
|
||||||
if len(args) == 0 {
|
|
||||||
if a.r.runtime.MaxSteps == 0 {
|
|
||||||
a.emit(infoEvent("maxsteps: unlimited"))
|
|
||||||
} else {
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("maxsteps: %d", a.r.runtime.MaxSteps)))
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
n, err := strconv.Atoi(args[0])
|
|
||||||
if err != nil || n < 0 {
|
|
||||||
a.emit(infoEvent("error: maxsteps requires a non-negative integer (0 = unlimited)"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
a.r.runtime.MaxSteps = n
|
|
||||||
if n == 0 {
|
|
||||||
a.emit(infoEvent("maxsteps: unlimited"))
|
|
||||||
} else {
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("maxsteps: %d", n)))
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
"/agents": func(args []string) {
|
|
||||||
seen := make(map[string]bool)
|
|
||||||
found := false
|
|
||||||
for _, dir := range AgentsDirs() {
|
|
||||||
entries, err := os.ReadDir(dir)
|
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
for _, e := range entries {
|
|
||||||
if e.IsDir() || !strings.HasSuffix(e.Name(), ".json") {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
name := strings.TrimSuffix(e.Name(), ".json")
|
|
||||||
if seen[name] {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
seen[name] = true
|
|
||||||
marker := " "
|
|
||||||
if name == a.r.agentName {
|
|
||||||
marker = "* "
|
|
||||||
}
|
|
||||||
a.emit(infoEvent(marker + name))
|
|
||||||
found = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
a.emit(infoEvent("no agents found"))
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
"/agent": func(args []string) {
|
|
||||||
if len(args) == 0 {
|
|
||||||
a.emit(infoEvent("active agent: " + a.r.agentName))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if a.IsRunning() {
|
|
||||||
a.emit(infoEvent("error: cannot switch agent while agent is running"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
name := args[0]
|
|
||||||
cfgPath := AgentConfigPath(a.r.agentsDir, name)
|
|
||||||
f, err := os.Open(cfgPath)
|
|
||||||
if err != nil {
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err)))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cfg, err := Load(f)
|
|
||||||
f.Close()
|
|
||||||
if err != nil {
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err)))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
d := a.r.newDispatcher()
|
|
||||||
env := []string{"OLLIE_SESSION_ID=" + a.id, "OLLIE_UNAME=" + a.r.id}
|
|
||||||
env = append(env, a.r.promptEnvExtra...)
|
|
||||||
rt := BuildRuntime(cfg, d, a.r.cwd, env, a.r.baseLayers...)
|
|
||||||
if rt.CfgBackend != "" {
|
|
||||||
newBe, err := a.r.newBackend(rt.CfgBackend)
|
|
||||||
if err != nil {
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("error: backend %q: %v", rt.CfgBackend, err)))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if rt.CfgModel != "" {
|
|
||||||
newBe.SetModel(rt.CfgModel)
|
|
||||||
}
|
|
||||||
rt.Backend = newBe
|
|
||||||
} else {
|
|
||||||
rt.Backend = a.r.runtime.Backend
|
|
||||||
if rt.CfgModel != "" {
|
|
||||||
rt.Backend.SetModel(rt.CfgModel)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
a.r.runtime = rt
|
|
||||||
a.r.agentName = name
|
|
||||||
a.r.history = nil
|
|
||||||
a.pushSessionEnv()
|
|
||||||
a.r.notifyChange()
|
|
||||||
for _, msg := range rt.Messages {
|
|
||||||
a.emit(infoEvent(msg))
|
|
||||||
}
|
|
||||||
a.emit(infoEvent("agent: " + name))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/compact": func(args []string) {
|
|
||||||
if a.IsRunning() {
|
|
||||||
a.emit(infoEvent("error: cannot compact while agent is running"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if a.r.history == nil {
|
|
||||||
a.emit(infoEvent("nothing to compact"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
snapshot := a.r.history.PreCompactionSnapshot()
|
|
||||||
a.r.SetState("compacting")
|
|
||||||
n, err := a.r.runCompact(ctx, "manual")
|
|
||||||
a.r.SetState("idle")
|
|
||||||
if err != nil {
|
|
||||||
a.emit(infoEvent("compact error: " + err.Error()))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if n == 0 {
|
|
||||||
a.emit(infoEvent("nothing to compact"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if a.sessionsDir != "" && a.id != "" {
|
|
||||||
histPath := a.activeSessionPath(a.id, ".compaction.jsonl")
|
|
||||||
if err := os.MkdirAll(filepath.Dir(histPath), 0700); err != nil {
|
|
||||||
a.emit(infoEvent("compaction history save: " + err.Error()))
|
|
||||||
} else if data, err := json.Marshal(snapshot); err == nil {
|
|
||||||
f, err := os.OpenFile(histPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0600)
|
|
||||||
if err == nil {
|
|
||||||
f.Write(append(data, '\n')) //nolint:errcheck
|
|
||||||
f.Close() //nolint:errcheck
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("compacted %d messages", n)))
|
|
||||||
a.saveSession()
|
|
||||||
},
|
|
||||||
|
|
||||||
"/context": func(args []string) {
|
|
||||||
if a.r.history == nil {
|
|
||||||
a.emit(infoEvent("no active session"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
ctxLen := a.r.runtime.Backend.ContextLength(ctx)
|
|
||||||
if ctxLen <= 0 {
|
|
||||||
ctxLen = defaultContextLength
|
|
||||||
}
|
|
||||||
estimated := a.r.history.estimateTokens()
|
|
||||||
pct := estimated * 100 / ctxLen
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("~%d / %d tokens (%d%%)", estimated, ctxLen, pct)))
|
|
||||||
a.emit(infoEvent(strings.TrimRight(a.r.history.contextDebug(), "\n")))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/cost": func(args []string) {
|
|
||||||
if a.r.history == nil {
|
|
||||||
a.emit(infoEvent("no active session"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("last=$%.4f session=$%.4f",
|
|
||||||
a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD)))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/usage": func(args []string) {
|
|
||||||
if a.r.history == nil {
|
|
||||||
a.emit(infoEvent("no active session"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
ctxLen := a.r.runtime.Backend.ContextLength(ctx)
|
|
||||||
if ctxLen <= 0 {
|
|
||||||
ctxLen = defaultContextLength
|
|
||||||
}
|
|
||||||
estimated := a.r.history.estimateTokens()
|
|
||||||
pct := estimated * 100 / ctxLen
|
|
||||||
usageStr := fmt.Sprintf("~%d / %d tokens (%d%%) | %d in, %d out, %d requests",
|
|
||||||
estimated, ctxLen, pct,
|
|
||||||
a.r.history.TotalInputTokens, a.r.history.TotalOutputTokens,
|
|
||||||
a.r.history.TotalRequests)
|
|
||||||
if a.r.history.Estimated {
|
|
||||||
usageStr += " [estimated]"
|
|
||||||
}
|
|
||||||
a.emit(infoEvent(usageStr))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/history": func(args []string) {
|
|
||||||
if a.r.history == nil {
|
|
||||||
a.emit(infoEvent("no active session"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
for _, msg := range a.r.history.history() {
|
|
||||||
preview := msg.Content
|
|
||||||
if len(preview) > 200 {
|
|
||||||
preview = preview[:200] + "..."
|
|
||||||
}
|
|
||||||
a.emit(infoEvent(fmt.Sprintf("[%s] %s", msg.Role, preview)))
|
|
||||||
}
|
|
||||||
},
|
|
||||||
|
|
||||||
"/clear": func(args []string) {
|
|
||||||
if a.IsRunning() {
|
|
||||||
a.emit(infoEvent("error: cannot clear while agent is running"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
a.r.history = nil
|
|
||||||
a.emit(infoEvent("cleared"))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/sessions": func(args []string) {
|
"/sessions": func(args []string) {
|
||||||
allFlag := len(args) > 0 && args[0] == "-a"
|
allFlag := len(args) > 0 && args[0] == "-a"
|
||||||
type sessionFile struct {
|
type sessionFile struct {
|
||||||
|
|
@ -343,7 +71,7 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
|
||||||
if readErr != nil {
|
if readErr != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
var ps PersistedAgent
|
var ps agent.PersistedAgent
|
||||||
if json.Unmarshal(data, &ps) != nil {
|
if json.Unmarshal(data, &ps) != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
@ -364,77 +92,73 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
|
||||||
if len(goal) > 60 {
|
if len(goal) > 60 {
|
||||||
goal = goal[:60] + "..."
|
goal = goal[:60] + "..."
|
||||||
}
|
}
|
||||||
a.emit(infoEvent(marker + fmt.Sprintf("%-24s [%s] %q", file.id, ps.Agent, goal)))
|
a.emit(agent.Event{Role: "info", Content: marker + fmt.Sprintf("%-24s [%s] %q", file.id, ps.Agent, goal) + "\n"})
|
||||||
found = true
|
found = true
|
||||||
}
|
}
|
||||||
if !found {
|
if !found {
|
||||||
a.emit(infoEvent("no sessions for " + cwd))
|
a.emit(agent.Event{Role: "info", Content: "no sessions for " + cwd + "\n"})
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
||||||
"/save": func(args []string) {
|
"/save": func(args []string) {
|
||||||
if a.r.history == nil {
|
if !a.r.HasHistory() {
|
||||||
a.emit(infoEvent("error: no active session"))
|
a.emit(agent.Event{Role: "info", Content: "error: no active session\n"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if len(args) == 0 {
|
if len(args) == 0 {
|
||||||
a.emit(infoEvent("error: /save requires a name"))
|
a.emit(agent.Event{Role: "info", Content: "error: /save requires a name\n"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
name := args[0]
|
name := args[0]
|
||||||
path := a.sessionsDir + "/" + name + ".json"
|
path := a.sessionsDir + "/" + name + ".json"
|
||||||
if err := a.r.history.saveTo(path, name, a.r.agentName, a.CWD()); err != nil {
|
if err := a.r.SaveTo(path, name, a.CWD()); err != nil {
|
||||||
a.emit(infoEvent("error: " + err.Error()))
|
a.emit(agent.Event{Role: "info", Content: "error: " + err.Error() + "\n"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
a.emit(infoEvent("saved: " + path))
|
a.emit(agent.Event{Role: "info", Content: "saved: " + path + "\n"})
|
||||||
},
|
},
|
||||||
|
|
||||||
"/resume": func(args []string) {
|
"/resume": func(args []string) {
|
||||||
if len(args) == 0 {
|
if len(args) == 0 {
|
||||||
a.emit(infoEvent("error: /resume requires a session id or name"))
|
a.emit(agent.Event{Role: "info", Content: "error: /resume requires a session id or name\n"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if a.IsRunning() {
|
if a.IsRunning() {
|
||||||
a.emit(infoEvent("error: cannot resume while agent is running"))
|
a.emit(agent.Event{Role: "info", Content: "error: cannot resume while agent is running\n"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
name := args[0]
|
name := args[0]
|
||||||
path := a.sessionsDir + "/" + name + ".json"
|
path := a.sessionsDir + "/" + name + ".json"
|
||||||
data, err := os.ReadFile(path)
|
data, err := os.ReadFile(path)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
a.emit(infoEvent(fmt.Sprintf("error: %v", err)))
|
a.emit(agent.Event{Role: "info", Content: fmt.Sprintf("error: %v\n", err)})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
var ps PersistedAgent
|
var ps agent.PersistedAgent
|
||||||
if err := json.Unmarshal(data, &ps); err != nil {
|
if err := json.Unmarshal(data, &ps); err != nil {
|
||||||
a.emit(infoEvent(fmt.Sprintf("error: %v", err)))
|
a.emit(agent.Event{Role: "info", Content: fmt.Sprintf("error: %v\n", err)})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
a.r.history = RestoreHistory(&ps)
|
a.r.Restore(&ps)
|
||||||
a.emit(infoEvent(fmt.Sprintf("resumed session %s (%d messages)", name, len(ps.Messages))))
|
a.emit(agent.Event{Role: "info", Content: fmt.Sprintf("resumed session %s (%d messages)\n", name, len(ps.Messages))})
|
||||||
},
|
},
|
||||||
|
|
||||||
"/cwd": func(args []string) {
|
"/cwd": func(args []string) {
|
||||||
if len(args) == 0 {
|
if len(args) == 0 {
|
||||||
a.emit(infoEvent("cwd: " + a.CWD()))
|
a.emit(agent.Event{Role: "info", Content: "cwd: " + a.CWD() + "\n"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
dir := strings.Join(args, " ")
|
dir := strings.Join(args, " ")
|
||||||
if err := a.SetCWD(dir); err != nil {
|
if err := a.SetCWD(dir); err != nil {
|
||||||
a.emit(infoEvent("error: " + err.Error()))
|
a.emit(agent.Event{Role: "info", Content: "error: " + err.Error() + "\n"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
a.emit(infoEvent("cwd: " + dir))
|
a.emit(agent.Event{Role: "info", Content: "cwd: " + dir + "\n"})
|
||||||
},
|
},
|
||||||
|
|
||||||
"/skills": func(args []string) { listFromHandler("skills") },
|
"/skills": func(args []string) { listFromHandler("skills") },
|
||||||
"/tools": func(args []string) { listFromHandler("tools") },
|
"/tools": func(args []string) { listFromHandler("tools") },
|
||||||
|
|
||||||
"/sp": func(args []string) {
|
|
||||||
a.emit(infoEvent(a.r.runtime.Preamble))
|
|
||||||
},
|
|
||||||
|
|
||||||
"/help": func(args []string) {
|
"/help": func(args []string) {
|
||||||
lines := []string{
|
lines := []string{
|
||||||
"Available commands:",
|
"Available commands:",
|
||||||
|
|
@ -452,21 +176,17 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
|
||||||
" /cwd [path] - show or change working directory",
|
" /cwd [path] - show or change working directory",
|
||||||
" /i <prompt> - inject prompt into the running turn",
|
" /i <prompt> - inject prompt into the running turn",
|
||||||
" /irw <prompt> - rewrite the pending inject",
|
" /irw <prompt> - rewrite the pending inject",
|
||||||
" /queued [pop|clear] - manage queued prompts",
|
|
||||||
" /compact - summarize conversation and compact context",
|
" /compact - summarize conversation and compact context",
|
||||||
" /context - show context size and message breakdown",
|
" /context - show context size and message breakdown",
|
||||||
" /cost - show last turn and session cost",
|
" /cost - show last turn and session cost",
|
||||||
" /usage - show token usage and context percentage",
|
" /usage - show token usage and context percentage",
|
||||||
" /history - dump bounded message history",
|
" /history - dump bounded message history",
|
||||||
" /clear - clear session",
|
" /clear - clear session",
|
||||||
" /kill - kill session",
|
|
||||||
" /rn <name> - rename session",
|
|
||||||
" /sp - show rendered system prompt",
|
" /sp - show rendered system prompt",
|
||||||
" /help - show this help",
|
" /help - show this help",
|
||||||
" !<cmd> - run shell command",
|
|
||||||
}
|
}
|
||||||
for _, l := range lines {
|
for _, l := range lines {
|
||||||
a.emit(infoEvent(l))
|
a.emit(agent.Event{Role: "info", Content: l + "\n"})
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
@ -475,7 +195,7 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
|
||||||
if !ok {
|
if !ok {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
a.emit(infoEvent(""))
|
a.emit(agent.Event{Role: "info", Content: "\n"})
|
||||||
fn(args)
|
fn(args)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,113 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
type mockBackendForCompact struct {
|
|
||||||
name string
|
|
||||||
model string
|
|
||||||
}
|
|
||||||
|
|
||||||
func (b *mockBackendForCompact) Name() string { return b.name }
|
|
||||||
func (b *mockBackendForCompact) Model() string { return b.model }
|
|
||||||
|
|
||||||
func TestResolveCompactionModel_ConfigWins(t *testing.T) {
|
|
||||||
b := &mockBackendForCompact{name: "anthropic", model: "claude-sonnet-4-5"}
|
|
||||||
got := resolveCompactionModel("my-custom-model", b)
|
|
||||||
if got != "my-custom-model" {
|
|
||||||
t.Errorf("got %q; want my-custom-model", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestResolveCompactionModel_EnvOverridesDefault(t *testing.T) {
|
|
||||||
t.Setenv("OLLIE_COMPACTION_MODEL", "env-model")
|
|
||||||
b := &mockBackendForCompact{name: "anthropic", model: "claude-sonnet-4-5"}
|
|
||||||
got := resolveCompactionModel("", b)
|
|
||||||
if got != "env-model" {
|
|
||||||
t.Errorf("got %q; want env-model", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestResolveCompactionModel_BackendDefault(t *testing.T) {
|
|
||||||
os.Unsetenv("OLLIE_COMPACTION_MODEL")
|
|
||||||
tests := []struct {
|
|
||||||
backend string
|
|
||||||
want string
|
|
||||||
}{
|
|
||||||
{"anthropic", "claude-3-5-haiku-latest"},
|
|
||||||
{"openai", "gpt-4o-mini"},
|
|
||||||
{"openrouter", "deepseek-v4-flash"},
|
|
||||||
{"gemini", "gemini-2.0-flash"},
|
|
||||||
{"ollama", ""},
|
|
||||||
}
|
|
||||||
for _, tt := range tests {
|
|
||||||
b := &mockBackendForCompact{name: tt.backend}
|
|
||||||
got := resolveCompactionModel("", b)
|
|
||||||
if got != tt.want {
|
|
||||||
t.Errorf("backend=%s: got %q; want %q", tt.backend, got, tt.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBuildCompactedHistory_OrphanedToolMessage(t *testing.T) {
|
|
||||||
// Simulate a history where the hot zone boundary (total - hotTailSize)
|
|
||||||
// lands on a tool message whose preceding assistant+tool_calls is outside
|
|
||||||
// the hot zone. Without the fix, this produces an orphaned tool message
|
|
||||||
// that causes OpenAI 400 errors.
|
|
||||||
var msgs []backend.Message
|
|
||||||
|
|
||||||
// Pad with enough messages so the boundary falls on the tool message.
|
|
||||||
// We need total - hotTailSize to land on the tool result.
|
|
||||||
// hotTailSize = 8, so we need the tool msg at index total-8.
|
|
||||||
// Build: 10 user/assistant pairs (20 msgs), then assistant+tool_calls, tool result, then 7 more messages.
|
|
||||||
for i := range 10 {
|
|
||||||
msgs = append(msgs,
|
|
||||||
backend.Message{Role: "user", Content: "q" + string(rune('0'+i))},
|
|
||||||
backend.Message{Role: "assistant", Content: "a" + string(rune('0'+i))},
|
|
||||||
)
|
|
||||||
}
|
|
||||||
// assistant with tool_calls at index 20
|
|
||||||
msgs = append(msgs, backend.Message{
|
|
||||||
Role: "assistant",
|
|
||||||
ToolCalls: []backend.ToolCall{{ID: "call_orphan", Name: "test_tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
})
|
|
||||||
// tool result at index 21 — this is where hotStart would land without the fix
|
|
||||||
msgs = append(msgs, backend.Message{
|
|
||||||
Role: "tool",
|
|
||||||
Content: "tool output",
|
|
||||||
ToolCallID: "call_orphan",
|
|
||||||
})
|
|
||||||
// 7 more messages to fill the rest of the hot zone (indices 22-28)
|
|
||||||
for i := range 3 {
|
|
||||||
msgs = append(msgs,
|
|
||||||
backend.Message{Role: "user", Content: "follow " + string(rune('0'+i))},
|
|
||||||
backend.Message{Role: "assistant", Content: "reply " + string(rune('0'+i))},
|
|
||||||
)
|
|
||||||
}
|
|
||||||
msgs = append(msgs, backend.Message{Role: "user", Content: "final"})
|
|
||||||
// total = 29, hotStart = 29 - 8 = 21 (the tool message)
|
|
||||||
|
|
||||||
ts := TaskState{Objective: "test"}
|
|
||||||
result := buildCompactedHistory(ts, msgs)
|
|
||||||
|
|
||||||
// Verify no tool message appears without a preceding assistant+tool_calls.
|
|
||||||
for i, m := range result {
|
|
||||||
if m.Role == "tool" {
|
|
||||||
if i == 0 {
|
|
||||||
t.Fatalf("result[0] is a tool message — no preceding assistant")
|
|
||||||
}
|
|
||||||
prev := result[i-1]
|
|
||||||
if prev.Role != "assistant" && prev.Role != "tool" {
|
|
||||||
t.Fatalf("result[%d] is tool but result[%d] is %q (want assistant or tool)", i, i-1, prev.Role)
|
|
||||||
}
|
|
||||||
if prev.Role == "assistant" && len(prev.ToolCalls) == 0 {
|
|
||||||
t.Fatalf("result[%d] is tool but preceding assistant has no tool_calls", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,137 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestLoad(t *testing.T) {
|
|
||||||
r := strings.NewReader(`{"hooks": {"postTurn": "notify-send done"}}`)
|
|
||||||
cfg, err := Load(r)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Load failed: %v", err)
|
|
||||||
}
|
|
||||||
if len(cfg.Hooks["postTurn"]) != 1 || cfg.Hooks["postTurn"][0] != "notify-send done" {
|
|
||||||
t.Errorf("Expected hook 'notify-send done', got %q", cfg.Hooks["postTurn"])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLoadEmpty(t *testing.T) {
|
|
||||||
cfg, err := Load(strings.NewReader(`{}`))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Load failed: %v", err)
|
|
||||||
}
|
|
||||||
if len(cfg.Hooks) != 0 {
|
|
||||||
t.Errorf("Expected no hooks, got %v", cfg.Hooks)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLoadInvalidJSON(t *testing.T) {
|
|
||||||
_, err := Load(strings.NewReader(`{bad`))
|
|
||||||
if err == nil {
|
|
||||||
t.Error("expected error for invalid JSON")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHookCmdsString(t *testing.T) {
|
|
||||||
cfg, err := Load(strings.NewReader(`{"hooks": {"pre": "single"}}`))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(cfg.Hooks["pre"]) != 1 || cfg.Hooks["pre"][0] != "single" {
|
|
||||||
t.Errorf("got %v, want [single]", cfg.Hooks["pre"])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHookCmdsArray(t *testing.T) {
|
|
||||||
cfg, err := Load(strings.NewReader(`{"hooks": {"pre": ["a", "b"]}}`))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(cfg.Hooks["pre"]) != 2 || cfg.Hooks["pre"][0] != "a" || cfg.Hooks["pre"][1] != "b" {
|
|
||||||
t.Errorf("got %v, want [a b]", cfg.Hooks["pre"])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHookCmdsInvalid(t *testing.T) {
|
|
||||||
_, err := Load(strings.NewReader(`{"hooks": {"pre": 42}}`))
|
|
||||||
if err == nil {
|
|
||||||
t.Error("expected error for invalid hook type")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPromptString(t *testing.T) {
|
|
||||||
cfg, err := Load(strings.NewReader(`{"prompt": "be helpful"}`))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if cfg.Prompt.IsExec || len(cfg.Prompt.Value) != 1 || cfg.Prompt.Value[0] != "be helpful" {
|
|
||||||
t.Errorf("Prompt = %v", cfg.Prompt)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPromptArray(t *testing.T) {
|
|
||||||
cfg, err := Load(strings.NewReader(`{"prompt": ["echo hello", "echo world"]}`))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if !cfg.Prompt.IsExec {
|
|
||||||
t.Error("expected IsExec=true for array prompt")
|
|
||||||
}
|
|
||||||
if len(cfg.Prompt.Value) != 2 || cfg.Prompt.Value[0] != "echo hello" || cfg.Prompt.Value[1] != "echo world" {
|
|
||||||
t.Errorf("Prompt.Value = %v", cfg.Prompt.Value)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPromptInvalid(t *testing.T) {
|
|
||||||
_, err := Load(strings.NewReader(`{"prompt": 42}`))
|
|
||||||
if err == nil {
|
|
||||||
t.Error("expected error for invalid prompt type")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLoadAllFields(t *testing.T) {
|
|
||||||
cfg, err := Load(strings.NewReader(`{
|
|
||||||
"prompt": "be helpful",
|
|
||||||
"maxTokens": 4096,
|
|
||||||
"temperature": 0.7,
|
|
||||||
"frequencyPenalty": 0.5,
|
|
||||||
"presencePenalty": 0.3
|
|
||||||
}`))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if cfg.Prompt.IsExec || len(cfg.Prompt.Value) != 1 || cfg.Prompt.Value[0] != "be helpful" {
|
|
||||||
t.Errorf("Prompt = %v", cfg.Prompt)
|
|
||||||
}
|
|
||||||
if cfg.MaxTokens != 4096 {
|
|
||||||
t.Errorf("MaxTokens = %d", cfg.MaxTokens)
|
|
||||||
}
|
|
||||||
if cfg.Temperature == nil || *cfg.Temperature != 0.7 {
|
|
||||||
t.Errorf("Temperature = %v", cfg.Temperature)
|
|
||||||
}
|
|
||||||
if cfg.FrequencyPenalty == nil || *cfg.FrequencyPenalty != 0.5 {
|
|
||||||
t.Errorf("FrequencyPenalty = %v", cfg.FrequencyPenalty)
|
|
||||||
}
|
|
||||||
if cfg.PresencePenalty == nil || *cfg.PresencePenalty != 0.3 {
|
|
||||||
t.Errorf("PresencePenalty = %v", cfg.PresencePenalty)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestToolsEnabled(t *testing.T) {
|
|
||||||
// Omitted: defaults to true.
|
|
||||||
cfg, _ := Load(strings.NewReader(`{}`))
|
|
||||||
if !cfg.ToolsEnabled() {
|
|
||||||
t.Error("expected ToolsEnabled()=true when omitted")
|
|
||||||
}
|
|
||||||
// Explicit false.
|
|
||||||
cfg, _ = Load(strings.NewReader(`{"tools": false}`))
|
|
||||||
if cfg.ToolsEnabled() {
|
|
||||||
t.Error("expected ToolsEnabled()=false")
|
|
||||||
}
|
|
||||||
// Explicit true.
|
|
||||||
cfg, _ = Load(strings.NewReader(`{"tools": true}`))
|
|
||||||
if !cfg.ToolsEnabled() {
|
|
||||||
t.Error("expected ToolsEnabled()=true")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
2595
session/core_test.go
2595
session/core_test.go
File diff suppressed because it is too large
Load Diff
|
|
@ -1,100 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"math"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
func approxEqual(a, b, tol float64) bool {
|
|
||||||
return math.Abs(a-b) <= tol
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestComputeCostUSD_BaseTokens verifies input+output pricing with no cache fields.
|
|
||||||
func TestComputeCostUSD_BaseTokens(t *testing.T) {
|
|
||||||
// claude-sonnet-4: $3/M input, $15/M output
|
|
||||||
got := computeCostUSD("claude-sonnet-4-5", backend.Usage{
|
|
||||||
InputTokens: 1_000_000,
|
|
||||||
OutputTokens: 1_000_000,
|
|
||||||
})
|
|
||||||
want := 3.00 + 15.00
|
|
||||||
if !approxEqual(got, want, 0.001) {
|
|
||||||
t.Errorf("got %.4f; want %.4f", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestComputeCostUSD_ClaudeCacheReadDiscount verifies that cached input tokens
|
|
||||||
// are charged at 10% of the normal input rate for Claude models.
|
|
||||||
func TestComputeCostUSD_ClaudeCacheReadDiscount(t *testing.T) {
|
|
||||||
// claude-sonnet-4: $3/M input → cache read = $0.30/M
|
|
||||||
got := computeCostUSD("claude-sonnet-4-5", backend.Usage{
|
|
||||||
InputTokens: 0,
|
|
||||||
CachedInputTokens: 1_000_000,
|
|
||||||
OutputTokens: 0,
|
|
||||||
})
|
|
||||||
want := 0.30 // 10% of $3.00
|
|
||||||
if !approxEqual(got, want, 0.001) {
|
|
||||||
t.Errorf("claude cache read: got %.4f; want %.4f", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestComputeCostUSD_NonClaudeCacheReadDiscount verifies that cached input tokens
|
|
||||||
// are charged at 50% of the normal input rate for non-Claude models (OpenAI).
|
|
||||||
func TestComputeCostUSD_NonClaudeCacheReadDiscount(t *testing.T) {
|
|
||||||
// gpt-4o: $2.50/M input → cache read = $1.25/M
|
|
||||||
got := computeCostUSD("gpt-4o", backend.Usage{
|
|
||||||
InputTokens: 0,
|
|
||||||
CachedInputTokens: 1_000_000,
|
|
||||||
OutputTokens: 0,
|
|
||||||
})
|
|
||||||
want := 1.25 // 50% of $2.50
|
|
||||||
if !approxEqual(got, want, 0.001) {
|
|
||||||
t.Errorf("gpt-4o cache read: got %.4f; want %.4f", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestComputeCostUSD_CacheCreationSurcharge verifies cache creation tokens are
|
|
||||||
// charged at 125% of the normal input rate.
|
|
||||||
func TestComputeCostUSD_CacheCreationSurcharge(t *testing.T) {
|
|
||||||
// claude-sonnet-4: $3/M input → cache creation = $3.75/M
|
|
||||||
got := computeCostUSD("claude-sonnet-4-5", backend.Usage{
|
|
||||||
InputTokens: 0,
|
|
||||||
CacheCreationTokens: 1_000_000,
|
|
||||||
OutputTokens: 0,
|
|
||||||
})
|
|
||||||
want := 3.75 // 125% of $3.00
|
|
||||||
if !approxEqual(got, want, 0.001) {
|
|
||||||
t.Errorf("cache creation surcharge: got %.4f; want %.4f", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestComputeCostUSD_AllFieldsCombined verifies that all four token categories
|
|
||||||
// are summed correctly in a single call.
|
|
||||||
func TestComputeCostUSD_AllFieldsCombined(t *testing.T) {
|
|
||||||
// claude-sonnet-4: $3/M in, $15/M out
|
|
||||||
// 100k normal input = $0.30
|
|
||||||
// 200k cached input = $0.06 (10% of $3/M)
|
|
||||||
// 50k cache create = $0.1875 (125% of $3/M)
|
|
||||||
// 100k output = $1.50
|
|
||||||
got := computeCostUSD("claude-sonnet-4-5", backend.Usage{
|
|
||||||
InputTokens: 100_000,
|
|
||||||
CachedInputTokens: 200_000,
|
|
||||||
CacheCreationTokens: 50_000,
|
|
||||||
OutputTokens: 100_000,
|
|
||||||
})
|
|
||||||
want := 0.30 + 0.06 + 0.1875 + 1.50
|
|
||||||
if !approxEqual(got, want, 0.0001) {
|
|
||||||
t.Errorf("combined: got %.6f; want %.6f", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestComputeCostUSD_UnknownModelZero verifies that unknown/local models return 0.
|
|
||||||
func TestComputeCostUSD_UnknownModelZero(t *testing.T) {
|
|
||||||
got := computeCostUSD("llama-3-local", backend.Usage{
|
|
||||||
InputTokens: 1_000_000, OutputTokens: 1_000_000,
|
|
||||||
})
|
|
||||||
if got != 0 {
|
|
||||||
t.Errorf("unknown model: got %.4f; want 0", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,187 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestResultCache_HitSkipsExec verifies that a second call to a read-safe tool
|
|
||||||
// with identical arguments returns the cached result without calling Exec again.
|
|
||||||
func TestResultCache_HitSkipsExec(t *testing.T) {
|
|
||||||
var execCount int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)},
|
|
||||||
{ID: "2", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(string) bool { return true }
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
atomic.AddInt32(&execCount, 1)
|
|
||||||
return "contents of a.txt", nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "dup read")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&execCount); n != 1 {
|
|
||||||
t.Errorf("Exec called %d times; want 1 (second should be cached)", n)
|
|
||||||
}
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 2 {
|
|
||||||
t.Errorf("tool events = %d; want 2", len(toolEvs))
|
|
||||||
}
|
|
||||||
for i, ev := range toolEvs {
|
|
||||||
if ev != "contents of a.txt" {
|
|
||||||
t.Errorf("tool event[%d] = %q; want cached value", i, ev)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestResultCache_DifferentArgsMiss verifies that the same tool name with
|
|
||||||
// different arguments produces two separate Exec calls (no false cache hits).
|
|
||||||
func TestResultCache_DifferentArgsMiss(t *testing.T) {
|
|
||||||
var execCount int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)},
|
|
||||||
{ID: "2", Name: "file_read", Arguments: json.RawMessage(`{"path":"b.txt"}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(string) bool { return true }
|
|
||||||
c.r.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
|
|
||||||
}
|
|
||||||
|
|
||||||
collectEvents(context.Background(), c, "different args")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&execCount); n != 2 {
|
|
||||||
t.Errorf("Exec called %d times; want 2 (different paths)", n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestResultCache_SerialToolNotCached verifies that non-read-safe tools are
|
|
||||||
// never cached: two identical calls both hit Exec.
|
|
||||||
func TestResultCache_SerialToolNotCached(t *testing.T) {
|
|
||||||
var execCount int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "file_write", Arguments: json.RawMessage(`{"path":"a.txt","content":"x"}`)},
|
|
||||||
{ID: "2", Name: "file_write", Arguments: json.RawMessage(`{"path":"a.txt","content":"x"}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(name string) bool { return false } // all serial
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
atomic.AddInt32(&execCount, 1)
|
|
||||||
return "ok", nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
collectEvents(context.Background(), c, "serial tool")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&execCount); n != 2 {
|
|
||||||
t.Errorf("Exec called %d times; want 2 (serial tools not cached)", n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// multiTurnToolsStream returns a backend respond function that issues a
|
|
||||||
// different set of tool calls on each successive invocation, then returns
|
|
||||||
// a plain text response once all sets are exhausted.
|
|
||||||
func multiTurnToolsStream(turns [][]backend.ToolCall) func(context.Context, []backend.Message, []backend.Tool, backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
var n int32
|
|
||||||
return func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
i := int(atomic.AddInt32(&n, 1)) - 1
|
|
||||||
if i < len(turns) {
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{ToolCalls: turns[i], Done: true, StopReason: "tool_calls"}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
return textStream("done"), nil
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestResultCache_ErrorNotCached verifies that a failed Exec result is not
|
|
||||||
// stored: a second call with the same args retries Exec rather than returning
|
|
||||||
// the cached error. The two calls are issued in separate loop turns so they
|
|
||||||
// run sequentially (no batching).
|
|
||||||
func TestResultCache_ErrorNotCached(t *testing.T) {
|
|
||||||
var execCount int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = multiTurnToolsStream([][]backend.ToolCall{
|
|
||||||
{{ID: "1", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)}},
|
|
||||||
{{ID: "2", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)}},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(string) bool { return true }
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
n := atomic.AddInt32(&execCount, 1)
|
|
||||||
if n == 1 {
|
|
||||||
return "", nil, &mockErr{"transient failure"}
|
|
||||||
}
|
|
||||||
return "ok now", nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "error then ok")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&execCount); n != 2 {
|
|
||||||
t.Errorf("Exec called %d times; want 2 (error must not be cached)", n)
|
|
||||||
}
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 2 {
|
|
||||||
t.Fatalf("tool events = %d; want 2", len(toolEvs))
|
|
||||||
}
|
|
||||||
if toolEvs[1] != "ok now" {
|
|
||||||
t.Errorf("second result = %q; want %q", toolEvs[1], "ok now")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestResultCache_IdenticalParallelReads verifies that two identical read-safe
|
|
||||||
// tool calls issued in the same parallel batch (both goroutines miss the cache
|
|
||||||
// simultaneously) both complete without error and return consistent results.
|
|
||||||
func TestResultCache_IdenticalParallelReads(t *testing.T) {
|
|
||||||
var execCount int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)},
|
|
||||||
{ID: "2", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(string) bool { return true } // both parallel-safe → same batch
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
atomic.AddInt32(&execCount, 1)
|
|
||||||
return "file contents", nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "identical parallel")
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 2 {
|
|
||||||
t.Fatalf("tool events = %d; want 2", len(toolEvs))
|
|
||||||
}
|
|
||||||
for i, ev := range toolEvs {
|
|
||||||
if ev != "file contents" {
|
|
||||||
t.Errorf("tool event[%d] = %q; want %q", i, ev, "file contents")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// Both may have executed (cache miss race) or one may have hit cache —
|
|
||||||
// either is correct. What must not happen: panic, empty result, or wrong value.
|
|
||||||
n := atomic.LoadInt32(&execCount)
|
|
||||||
if n < 1 || n > 2 {
|
|
||||||
t.Errorf("Exec called %d times; want 1 or 2", n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type mockErr struct{ msg string }
|
|
||||||
|
|
||||||
func (e *mockErr) Error() string { return e.msg }
|
|
||||||
|
|
@ -1,178 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"strings"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
// alwaysFailStream returns a backend that issues a single tool call on every
|
|
||||||
// invocation, never producing a text-only (stop) response.
|
|
||||||
func alwaysFailStream() func(context.Context, []backend.Message, []backend.Tool, backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
return func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{
|
|
||||||
ToolCalls: []backend.ToolCall{{ID: "c1", Name: "bad_tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
Done: true,
|
|
||||||
StopReason: "tool_calls",
|
|
||||||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestConsecutiveErrors_HardLimit(t *testing.T) {
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = alwaysFailStream()
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return "", nil, fmt.Errorf("always fails")
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "do something")
|
|
||||||
|
|
||||||
// Should see an error event about consecutive tool errors.
|
|
||||||
errs := byRole(evs, "error")
|
|
||||||
found := false
|
|
||||||
for _, e := range errs {
|
|
||||||
if strings.Contains(e, "consecutive tool errors") {
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
t.Errorf("expected 'consecutive tool errors' in error events; got %v", errs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestConsecutiveErrors_SoftLimitNudge(t *testing.T) {
|
|
||||||
var rounds atomic.Int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = func(_ context.Context, msgs []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
n := int(rounds.Add(1))
|
|
||||||
// After soft limit, check that the nudge was injected into the
|
|
||||||
// conversation history. Stop the loop by returning text only.
|
|
||||||
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
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// Nudge not found — keep going (will hit hard limit if broken).
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{
|
|
||||||
ToolCalls: []backend.ToolCall{{ID: "c1", Name: "bad_tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
Done: true,
|
|
||||||
StopReason: "tool_calls",
|
|
||||||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{
|
|
||||||
ToolCalls: []backend.ToolCall{{ID: "c1", Name: "bad_tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
Done: true,
|
|
||||||
StopReason: "tool_calls",
|
|
||||||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return "", nil, fmt.Errorf("always fails")
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "do something")
|
|
||||||
|
|
||||||
// The model should have seen the nudge and responded with text, ending the loop.
|
|
||||||
texts := byRole(evs, "assistant")
|
|
||||||
found := false
|
|
||||||
for _, txt := range texts {
|
|
||||||
if strings.Contains(txt, "giving up") {
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
t.Error("expected model to receive nudge and respond with 'giving up'")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestConsecutiveErrors_ResetOnSuccess(t *testing.T) {
|
|
||||||
var rounds atomic.Int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
n := int(rounds.Add(1))
|
|
||||||
// Rounds 1-4: fail. Round 5: succeed. Rounds 6-9: fail. Round 10: succeed. Round 11: text.
|
|
||||||
// This ensures the counter resets and we never hit the soft limit.
|
|
||||||
if n == 5 || n == 10 {
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{
|
|
||||||
ToolCalls: []backend.ToolCall{{ID: "c1", Name: "good_tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
Done: true,
|
|
||||||
StopReason: "tool_calls",
|
|
||||||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
if n >= 11 {
|
|
||||||
return textStream("done"), nil
|
|
||||||
|
|
||||||
}
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{
|
|
||||||
ToolCalls: []backend.ToolCall{{ID: "c1", Name: "bad_tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
Done: true,
|
|
||||||
StopReason: "tool_calls",
|
|
||||||
Usage: backend.Usage{InputTokens: 10, OutputTokens: 5},
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
if name == "good_tool" {
|
|
||||||
return "ok", nil, nil
|
|
||||||
}
|
|
||||||
return "", nil, fmt.Errorf("fails")
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "do something")
|
|
||||||
|
|
||||||
// Should complete normally — no consecutive-error abort.
|
|
||||||
errs := byRole(evs, "error")
|
|
||||||
for _, e := range errs {
|
|
||||||
if strings.Contains(e, "consecutive tool errors") {
|
|
||||||
t.Errorf("unexpected hard limit error; counter should have reset: %s", e)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
texts := byRole(evs, "assistant")
|
|
||||||
found := false
|
|
||||||
for _, txt := range texts {
|
|
||||||
if strings.Contains(txt, "done") {
|
|
||||||
found = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !found {
|
|
||||||
t.Error("expected loop to complete normally with 'done' response")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,76 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"sync/atomic"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
// mockResponse defines a canned response for sequentialStream.
|
|
||||||
type mockResponse struct {
|
|
||||||
content string
|
|
||||||
toolCalls []backend.ToolCall
|
|
||||||
stopReason string
|
|
||||||
}
|
|
||||||
|
|
||||||
// sequentialStream returns a respond function that plays back responses in order.
|
|
||||||
func sequentialStream(responses []mockResponse) func(context.Context, []backend.Message, []backend.Tool, backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
var n int32
|
|
||||||
return func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
i := int(atomic.AddInt32(&n, 1)) - 1
|
|
||||||
if i >= len(responses) {
|
|
||||||
return textStream("done"), nil
|
|
||||||
|
|
||||||
}
|
|
||||||
r := responses[i]
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{
|
|
||||||
Content: r.content,
|
|
||||||
ToolCalls: r.toolCalls,
|
|
||||||
Done: true,
|
|
||||||
StopReason: r.stopReason,
|
|
||||||
}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// simpleState is a minimal state implementation for direct run() calls.
|
|
||||||
type simpleState struct {
|
|
||||||
msgs []backend.Message
|
|
||||||
}
|
|
||||||
|
|
||||||
func newState() *simpleState {
|
|
||||||
return &simpleState{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *simpleState) history() []backend.Message {
|
|
||||||
return s.msgs
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *simpleState) taskState() *TaskState { return nil }
|
|
||||||
|
|
||||||
func (s *simpleState) updateTaskState(TaskState) {}
|
|
||||||
|
|
||||||
func (s *simpleState) update(msg backend.Message, results []toolResult) {
|
|
||||||
s.msgs = append(s.msgs, msg)
|
|
||||||
for _, r := range results {
|
|
||||||
s.msgs = append(s.msgs, backend.Message{
|
|
||||||
Role: "tool",
|
|
||||||
Content: r.Content,
|
|
||||||
ToolCallID: r.ToolCallID,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *simpleState) estimateTokens() int {
|
|
||||||
chars := 0
|
|
||||||
for _, m := range s.msgs {
|
|
||||||
chars += len(m.Content)
|
|
||||||
}
|
|
||||||
return chars / 4
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *simpleState) stripCold(_ context.Context, _ backend.Backend) {}
|
|
||||||
|
|
@ -1,184 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestMaxStepsZeroUnlimited verifies that MaxSteps=0 does not trigger the
|
|
||||||
// guardrail — the loop runs until the model stops calling tools.
|
|
||||||
func TestMaxStepsZeroUnlimited(t *testing.T) {
|
|
||||||
var steps int
|
|
||||||
mb := &mockBackend{
|
|
||||||
respond: sequentialStream([]mockResponse{
|
|
||||||
{toolCalls: []backend.ToolCall{{ID: "1", Name: "tool", Arguments: json.RawMessage(`{}`)}}, stopReason: "tool_calls"},
|
|
||||||
{toolCalls: []backend.ToolCall{{ID: "2", Name: "tool", Arguments: json.RawMessage(`{}`)}}, stopReason: "tool_calls"},
|
|
||||||
{content: "done", stopReason: "stop"},
|
|
||||||
}),
|
|
||||||
}
|
|
||||||
cfg := agentConfig{
|
|
||||||
Backend: mb,
|
|
||||||
Tools: []backend.Tool{{Name: "tool"}},
|
|
||||||
MaxSteps: 0,
|
|
||||||
Exec: func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
steps++
|
|
||||||
return "ok", nil, nil
|
|
||||||
},
|
|
||||||
}
|
|
||||||
if err := run(context.Background(), cfg, newState()); err != nil {
|
|
||||||
t.Fatalf("unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
if steps != 2 {
|
|
||||||
t.Errorf("expected 2 tool executions, got %d", steps)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMaxStepsSoftNudge verifies that when MaxSteps is reached the loop injects
|
|
||||||
// the budget-exhausted nudge message, emits a maxsteps event, and exits cleanly
|
|
||||||
// (no error returned). The model is given one final tool-free turn.
|
|
||||||
func TestMaxStepsSoftNudge(t *testing.T) {
|
|
||||||
var nudgeSeen bool
|
|
||||||
var maxstepsEventSeen bool
|
|
||||||
|
|
||||||
// Backend: two tool-calling rounds, then a final text turn.
|
|
||||||
mb := &mockBackend{
|
|
||||||
respond: sequentialStream([]mockResponse{
|
|
||||||
{toolCalls: []backend.ToolCall{{ID: "1", Name: "tool", Arguments: json.RawMessage(`{}`)}}, stopReason: "tool_calls"},
|
|
||||||
{toolCalls: []backend.ToolCall{{ID: "2", Name: "tool", Arguments: json.RawMessage(`{}`)}}, stopReason: "tool_calls"},
|
|
||||||
{content: "wrapping up", stopReason: "stop"},
|
|
||||||
}),
|
|
||||||
}
|
|
||||||
|
|
||||||
// MaxSteps=1 means the guardrail fires after completing step 0 (the first
|
|
||||||
// tool round), before step 1 would begin.
|
|
||||||
cfg := agentConfig{
|
|
||||||
Backend: mb,
|
|
||||||
Tools: []backend.Tool{{Name: "tool"}},
|
|
||||||
MaxSteps: 1,
|
|
||||||
Exec: func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return "ok", nil, nil
|
|
||||||
},
|
|
||||||
Output: func(ev Event) {
|
|
||||||
if ev.Role == "maxsteps" {
|
|
||||||
maxstepsEventSeen = true
|
|
||||||
}
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// Intercept state updates to detect the nudge message.
|
|
||||||
s := newState()
|
|
||||||
origUpdate := s.update
|
|
||||||
_ = origUpdate // state.update is not a field; we'll check history post-run instead.
|
|
||||||
|
|
||||||
if err := run(context.Background(), cfg, s); err != nil {
|
|
||||||
t.Fatalf("unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !maxstepsEventSeen {
|
|
||||||
t.Error("expected maxsteps event to be emitted")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Confirm the nudge message is present in conversation history.
|
|
||||||
for _, m := range s.history() {
|
|
||||||
if m.Role == "user" && strings.Contains(m.Content, "step budget exhausted") {
|
|
||||||
nudgeSeen = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !nudgeSeen {
|
|
||||||
t.Error("expected step-budget nudge message in conversation history")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestMaxStepsExactBoundary checks that with MaxSteps=N the loop completes
|
|
||||||
// exactly N tool rounds before nudging.
|
|
||||||
func TestMaxStepsExactBoundary(t *testing.T) {
|
|
||||||
var toolRounds int
|
|
||||||
|
|
||||||
const limit = 3
|
|
||||||
|
|
||||||
// Build limit+1 tool responses so the model would run forever without the cap.
|
|
||||||
var responses []mockResponse
|
|
||||||
for i := range limit + 1 {
|
|
||||||
responses = append(responses, mockResponse{
|
|
||||||
toolCalls: []backend.ToolCall{{ID: fmt.Sprintf("%d", i+1), Name: "tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
stopReason: "tool_calls",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
responses = append(responses, mockResponse{content: "done", stopReason: "stop"})
|
|
||||||
|
|
||||||
mb := &mockBackend{respond: sequentialStream(responses)}
|
|
||||||
cfg := agentConfig{
|
|
||||||
Backend: mb,
|
|
||||||
Tools: []backend.Tool{{Name: "tool"}},
|
|
||||||
MaxSteps: limit,
|
|
||||||
Exec: func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
toolRounds++
|
|
||||||
return "ok", nil, nil
|
|
||||||
},
|
|
||||||
}
|
|
||||||
if err := run(context.Background(), cfg, newState()); err != nil {
|
|
||||||
t.Fatalf("unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
if toolRounds != limit {
|
|
||||||
t.Errorf("expected %d tool rounds, got %d", limit, toolRounds)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTaskStateReinjection verifies that the loop re-injects the task state
|
|
||||||
// into the conversation every planReinjectInterval tool rounds.
|
|
||||||
func TestTaskStateReinjection(t *testing.T) {
|
|
||||||
// We need planReinjectInterval+1 tool rounds so the re-injection fires
|
|
||||||
// at step == planReinjectInterval (0-indexed, checked after increment).
|
|
||||||
n := planReinjectInterval + 1
|
|
||||||
var responses []mockResponse
|
|
||||||
for i := range n {
|
|
||||||
responses = append(responses, mockResponse{
|
|
||||||
toolCalls: []backend.ToolCall{{ID: fmt.Sprintf("%d", i+1), Name: "tool", Arguments: json.RawMessage(`{}`)}},
|
|
||||||
stopReason: "tool_calls",
|
|
||||||
})
|
|
||||||
}
|
|
||||||
responses = append(responses, mockResponse{content: "done", stopReason: "stop"})
|
|
||||||
|
|
||||||
mb := &mockBackend{respond: sequentialStream(responses)}
|
|
||||||
cfg := agentConfig{
|
|
||||||
Backend: mb,
|
|
||||||
Tools: []backend.Tool{{Name: "tool"}},
|
|
||||||
Exec: func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return "ok", nil, nil
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
s := &taskStateState{
|
|
||||||
simpleState: simpleState{},
|
|
||||||
ts: &TaskState{Objective: "test objective", PlanStep: "step one"},
|
|
||||||
}
|
|
||||||
if err := run(context.Background(), cfg, s); err != nil {
|
|
||||||
t.Fatalf("unexpected error: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var seen bool
|
|
||||||
for _, m := range s.history() {
|
|
||||||
if m.Role == "user" && strings.Contains(m.Content, "test objective") && strings.Contains(m.Content, "task state") {
|
|
||||||
seen = true
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if !seen {
|
|
||||||
t.Error("expected task state re-injection message in conversation history")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// taskStateState wraps simpleState with a non-nil TaskState.
|
|
||||||
type taskStateState struct {
|
|
||||||
simpleState
|
|
||||||
ts *TaskState
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *taskStateState) taskState() *TaskState { return s.ts }
|
|
||||||
func (s *taskStateState) updateTaskState(ts TaskState) { s.ts = &ts }
|
|
||||||
|
|
@ -1,206 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"sort"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
// toolsStream returns a backend respond function that issues the given tool
|
|
||||||
// calls on the first invocation and returns a plain text response thereafter.
|
|
||||||
func toolsStream(calls []backend.ToolCall) func(context.Context, []backend.Message, []backend.Tool, backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
var n int32
|
|
||||||
return func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
if atomic.AddInt32(&n, 1) == 1 {
|
|
||||||
ch := make(chan backend.StreamEvent, 1)
|
|
||||||
ch <- backend.StreamEvent{ToolCalls: calls, Done: true, StopReason: "tool_calls"}
|
|
||||||
close(ch)
|
|
||||||
return ch, nil
|
|
||||||
|
|
||||||
}
|
|
||||||
return textStream("done"), nil
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestParallel_ConcurrentExecution proves that read-safe tools actually run
|
|
||||||
// in parallel. The barrier requires all 3 goroutines to be in-flight at the
|
|
||||||
// same time; sequential execution would deadlock and trip the timeout.
|
|
||||||
func TestParallel_ConcurrentExecution(t *testing.T) {
|
|
||||||
const n = 3
|
|
||||||
started := make(chan struct{}, n)
|
|
||||||
gate := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
for i := 0; i < n; i++ {
|
|
||||||
<-started
|
|
||||||
}
|
|
||||||
close(gate) // open once all n tools have started
|
|
||||||
}()
|
|
||||||
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "tool_a", Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "2", Name: "tool_b", Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "3", Name: "tool_c", Arguments: json.RawMessage(`{}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(string) bool { return true }
|
|
||||||
c.r.runtime.Exec = func(ctx context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
started <- struct{}{}
|
|
||||||
select {
|
|
||||||
case <-gate:
|
|
||||||
return name + "-result", nil, nil
|
|
||||||
case <-ctx.Done():
|
|
||||||
return "", nil, ctx.Err()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
evs := collectEvents(ctx, c, "run parallel")
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
t.Fatal("timed out — tools likely ran sequentially (barrier never opened)")
|
|
||||||
}
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != n {
|
|
||||||
t.Errorf("tool events = %d; want %d", len(toolEvs), n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestParallel_SerialToolBreaksBatch verifies that a serial tool between two
|
|
||||||
// read-safe tools prevents them from being batched together. Order must be
|
|
||||||
// read_a, write_b, read_c regardless of internal execution details.
|
|
||||||
func TestParallel_SerialToolBreaksBatch(t *testing.T) {
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "read_a", Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "2", Name: "write_b", Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "3", Name: "read_c", Arguments: json.RawMessage(`{}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
var mu sync.Mutex
|
|
||||||
var order []string
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(name string) bool {
|
|
||||||
return name == "read_a" || name == "read_c"
|
|
||||||
}
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
mu.Lock()
|
|
||||||
order = append(order, name)
|
|
||||||
mu.Unlock()
|
|
||||||
return name + "-result", nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
collectEvents(context.Background(), c, "mixed tools")
|
|
||||||
|
|
||||||
mu.Lock()
|
|
||||||
got := append([]string(nil), order...)
|
|
||||||
mu.Unlock()
|
|
||||||
|
|
||||||
if len(got) != 3 {
|
|
||||||
t.Fatalf("executions = %d; want 3: %v", len(got), got)
|
|
||||||
}
|
|
||||||
// write_b must appear after the reads that precede it and before those that follow.
|
|
||||||
// With single-element batches for reads flanking a serial write, order is deterministic.
|
|
||||||
if got[0] != "read_a" || got[1] != "write_b" || got[2] != "read_c" {
|
|
||||||
t.Errorf("execution order = %v; want [read_a write_b read_c]", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestParallel_CancellationFillsRemaining verifies that when the context is
|
|
||||||
// cancelled during a parallel batch, all tools in the batch still produce
|
|
||||||
// results (IsError) and any subsequent tool calls also get cancelled results.
|
|
||||||
func TestParallel_CancellationFillsRemaining(t *testing.T) {
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "tool_a", Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "2", Name: "tool_b", Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "3", Name: "tool_c", Arguments: json.RawMessage(`{}`)}, // serial — after the parallel batch
|
|
||||||
})
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
// tool_a and tool_b are parallel-safe; tool_c is serial.
|
|
||||||
c.r.runtime.ClassifyTool = func(name string) bool {
|
|
||||||
return name == "tool_a" || name == "tool_b"
|
|
||||||
}
|
|
||||||
c.r.runtime.Exec = func(execCtx context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
cancel() // cancel on first execution; propagates to all
|
|
||||||
<-execCtx.Done()
|
|
||||||
return "", nil, execCtx.Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(ctx, c, "cancel mid-batch")
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
// All 3 tool calls must have produced a result.
|
|
||||||
if len(toolEvs) != 3 {
|
|
||||||
t.Errorf("tool events = %d; want 3", len(toolEvs))
|
|
||||||
}
|
|
||||||
// All results must be errors.
|
|
||||||
for _, ev := range toolEvs {
|
|
||||||
_ = ev // content varies; IsError is tracked internally, not in the event text
|
|
||||||
}
|
|
||||||
// The backend should not have been called a second time (interrupted before follow-up).
|
|
||||||
for _, ev := range evs {
|
|
||||||
if ev.Role == "assistant" && ev.Content == "done" {
|
|
||||||
t.Error("follow-up 'done' response received; expected interruption before second backend call")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestParallel_NilClassifyToolIsSerial verifies that when ClassifyTool is nil
|
|
||||||
// all tools execute sequentially and all results are returned in order.
|
|
||||||
func TestParallel_NilClassifyToolIsSerial(t *testing.T) {
|
|
||||||
names := []string{"tool_a", "tool_b", "tool_c"}
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: names[0], Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "2", Name: names[1], Arguments: json.RawMessage(`{}`)},
|
|
||||||
{ID: "3", Name: names[2], Arguments: json.RawMessage(`{}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
var mu sync.Mutex
|
|
||||||
var order []string
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = nil // no classifier → all serial
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, name string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
mu.Lock()
|
|
||||||
order = append(order, name)
|
|
||||||
mu.Unlock()
|
|
||||||
return name + "-result", nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "serial fallback")
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 3 {
|
|
||||||
t.Errorf("tool events = %d; want 3", len(toolEvs))
|
|
||||||
}
|
|
||||||
|
|
||||||
mu.Lock()
|
|
||||||
got := append([]string(nil), order...)
|
|
||||||
mu.Unlock()
|
|
||||||
|
|
||||||
sort.Strings(got)
|
|
||||||
sort.Strings(names)
|
|
||||||
for i, g := range got {
|
|
||||||
if g != names[i] {
|
|
||||||
t.Errorf("execution order mismatch: got %v", got)
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,126 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestTruncation_LargeResultTruncated verifies that a tool result exceeding
|
|
||||||
// the 128KB safety limit is truncated and a hint is appended.
|
|
||||||
func TestTruncation_LargeResultTruncated(t *testing.T) {
|
|
||||||
large := strings.Repeat("x", defaultToolResultMaxBytes+10_000)
|
|
||||||
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "execute_code", Arguments: json.RawMessage(`{}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return large, nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "big result")
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 1 {
|
|
||||||
t.Fatalf("tool events = %d; want 1", len(toolEvs))
|
|
||||||
}
|
|
||||||
result := toolEvs[0]
|
|
||||||
if !strings.HasPrefix(result, strings.Repeat("x", defaultToolResultMaxBytes)) {
|
|
||||||
t.Errorf("result does not start with %d x's", defaultToolResultMaxBytes)
|
|
||||||
}
|
|
||||||
if !strings.Contains(result, "HARD LIMIT") {
|
|
||||||
t.Errorf("truncation hint missing from result: %q", result[:min(len(result), 80)])
|
|
||||||
}
|
|
||||||
if !strings.Contains(result, "execute_code") {
|
|
||||||
t.Errorf("tool name missing from truncation hint: %q", result[:min(len(result), 120)])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTruncation_SmallResultNotTruncated verifies that results within the limit
|
|
||||||
// pass through unchanged.
|
|
||||||
func TestTruncation_SmallResultNotTruncated(t *testing.T) {
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "execute_code", Arguments: json.RawMessage(`{}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return "short result", nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "small result")
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 1 {
|
|
||||||
t.Fatalf("tool events = %d; want 1", len(toolEvs))
|
|
||||||
}
|
|
||||||
if toolEvs[0] != "short result" {
|
|
||||||
t.Errorf("result = %q; want %q", toolEvs[0], "short result")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTruncation_ErrorTruncated verifies that error results are subject to the
|
|
||||||
// same 128KB safety limit (the original bug: errors bypassed all truncation).
|
|
||||||
func TestTruncation_ErrorTruncated(t *testing.T) {
|
|
||||||
longErr := strings.Repeat("e", defaultToolResultMaxBytes+10_000)
|
|
||||||
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = toolsStream([]backend.ToolCall{
|
|
||||||
{ID: "1", Name: "execute_code", Arguments: json.RawMessage(`{}`)},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return "", nil, &mockErr{longErr}
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "error result")
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 1 {
|
|
||||||
t.Fatalf("tool events = %d; want 1", len(toolEvs))
|
|
||||||
}
|
|
||||||
if !strings.Contains(toolEvs[0], "HARD LIMIT") {
|
|
||||||
t.Errorf("error result was NOT truncated: len=%d", len(toolEvs[0]))
|
|
||||||
}
|
|
||||||
if len(toolEvs[0]) > defaultToolResultMaxBytes+200 {
|
|
||||||
t.Errorf("error result too large after truncation: len=%d", len(toolEvs[0]))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTruncation_CachedResultAlreadyTruncated verifies that a cache hit on a
|
|
||||||
// previously-truncated result returns the truncated form, not the original.
|
|
||||||
func TestTruncation_CachedResultAlreadyTruncated(t *testing.T) {
|
|
||||||
large := strings.Repeat("z", defaultToolResultMaxBytes+10_000)
|
|
||||||
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = multiTurnToolsStream([][]backend.ToolCall{
|
|
||||||
{{ID: "1", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)}},
|
|
||||||
{{ID: "2", Name: "file_read", Arguments: json.RawMessage(`{"path":"a.txt"}`)}},
|
|
||||||
})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.runtime.ClassifyTool = func(string) bool { return true } // read-safe → cacheable
|
|
||||||
c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) {
|
|
||||||
return large, nil, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "cached truncated")
|
|
||||||
|
|
||||||
toolEvs := byRole(evs, "tool")
|
|
||||||
if len(toolEvs) != 2 {
|
|
||||||
t.Fatalf("tool events = %d; want 2", len(toolEvs))
|
|
||||||
}
|
|
||||||
for i, ev := range toolEvs {
|
|
||||||
if !strings.Contains(ev, "HARD LIMIT") {
|
|
||||||
t.Errorf("event[%d] missing truncation hint: %q", i, ev[:min(len(ev), 80)])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,164 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"ollie/backend"
|
|
||||||
)
|
|
||||||
|
|
||||||
// errStream returns a backend respond function that always returns the given error.
|
|
||||||
func errStream(err error) func(context.Context, []backend.Message, []backend.Tool, backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
return func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// errThenOKStream returns an error on the first call, then a text response.
|
|
||||||
func errThenOKStream(err error) func(context.Context, []backend.Message, []backend.Tool, backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
var n int32
|
|
||||||
return func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
if atomic.AddInt32(&n, 1) == 1 {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return textStream("recovered"), nil
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTurnError_HookInterceptsRateLimit verifies that a turnError hook fired
|
|
||||||
// on a RateLimitError causes the loop to skip retries and return immediately.
|
|
||||||
func TestTurnError_HookInterceptsRateLimit(t *testing.T) {
|
|
||||||
var hookCalls int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = errStream(&backend.RateLimitError{Message: "quota exceeded"})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.turnError = func(_ context.Context, errType, _ string) HookResult {
|
|
||||||
atomic.AddInt32(&hookCalls, 1)
|
|
||||||
if errType != "rate_limit" {
|
|
||||||
t.Errorf("errType = %q; want rate_limit", errType)
|
|
||||||
}
|
|
||||||
return HookResult{Ran: true, Handled: true}
|
|
||||||
}
|
|
||||||
|
|
||||||
evs := collectEvents(context.Background(), c, "hi")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&hookCalls); n != 1 {
|
|
||||||
t.Errorf("hook called %d times; want 1 (no retries after hook intercept)", n)
|
|
||||||
}
|
|
||||||
errEvs := byRole(evs, "error")
|
|
||||||
if len(errEvs) == 0 {
|
|
||||||
t.Error("expected an error event")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTurnError_HookInterceptsToolUnsupported verifies the same skip-retry
|
|
||||||
// behaviour for ToolUnsupportedError.
|
|
||||||
func TestTurnError_HookInterceptsToolUnsupported(t *testing.T) {
|
|
||||||
var hookCalls int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = errStream(&backend.ToolUnsupportedError{Message: "model does not support tools"})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.turnError = func(_ context.Context, errType, _ string) HookResult {
|
|
||||||
atomic.AddInt32(&hookCalls, 1)
|
|
||||||
if errType != "tool_unsupported" {
|
|
||||||
t.Errorf("errType = %q; want tool_unsupported", errType)
|
|
||||||
}
|
|
||||||
return HookResult{Ran: true, Handled: true}
|
|
||||||
}
|
|
||||||
|
|
||||||
collectEvents(context.Background(), c, "hi")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&hookCalls); n != 1 {
|
|
||||||
t.Errorf("hook called %d times; want 1", n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTurnError_NoHookFallsThrough verifies that when no turnError hook is
|
|
||||||
// configured, normal retry behaviour proceeds for retryable errors.
|
|
||||||
func TestTurnError_NoHookFallsThrough(t *testing.T) {
|
|
||||||
old := retryBaseDelay
|
|
||||||
retryBaseDelay = 10 * time.Millisecond
|
|
||||||
defer func() { retryBaseDelay = old }()
|
|
||||||
|
|
||||||
var attempts int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) {
|
|
||||||
atomic.AddInt32(&attempts, 1)
|
|
||||||
return nil, &backend.RateLimitError{Message: "slow down"}
|
|
||||||
}
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
// TurnError is nil — no hook configured.
|
|
||||||
|
|
||||||
collectEvents(context.Background(), c, "hi")
|
|
||||||
|
|
||||||
// Should have attempted maxTransientRetries+1 = 4 times.
|
|
||||||
if n := atomic.LoadInt32(&attempts); n != maxTransientRetries+1 {
|
|
||||||
t.Errorf("attempts = %d; want %d (full retry cycle)", n, maxTransientRetries+1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTurnError_NonRetryableErrorNoHook verifies that a plain (non-retryable)
|
|
||||||
// error fires the hook once and does not retry.
|
|
||||||
func TestTurnError_NonRetryableErrorNoHook(t *testing.T) {
|
|
||||||
var hookCalls int32
|
|
||||||
be := defaultBE()
|
|
||||||
be.respond = errStream(&backend.ToolUnsupportedError{Message: "no tools"})
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.turnError = func(_ context.Context, errType, _ string) HookResult {
|
|
||||||
atomic.AddInt32(&hookCalls, 1)
|
|
||||||
return HookResult{Ran: true, Handled: true}
|
|
||||||
}
|
|
||||||
|
|
||||||
collectEvents(context.Background(), c, "hi")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&hookCalls); n != 1 {
|
|
||||||
t.Errorf("hook called %d times; want 1", n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTurnError_HookNotRunOnSuccess verifies that the turnError hook is never
|
|
||||||
// called when the backend succeeds on the first attempt.
|
|
||||||
func TestTurnError_HookNotRunOnSuccess(t *testing.T) {
|
|
||||||
var hookCalls int32
|
|
||||||
be := defaultBE()
|
|
||||||
// Default respond returns textStream("ok") — no error.
|
|
||||||
|
|
||||||
c := newCore(t, be, nil)
|
|
||||||
c.r.turnError = func(_ context.Context, _, _ string) HookResult {
|
|
||||||
atomic.AddInt32(&hookCalls, 1)
|
|
||||||
return HookResult{Ran: true}
|
|
||||||
}
|
|
||||||
|
|
||||||
collectEvents(context.Background(), c, "hi")
|
|
||||||
|
|
||||||
if n := atomic.LoadInt32(&hookCalls); n != 0 {
|
|
||||||
t.Errorf("hook called %d times on success; want 0", n)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestTurnError_ClassifyError verifies that classifyError returns the correct
|
|
||||||
// string for each known error type.
|
|
||||||
func TestTurnError_ClassifyError(t *testing.T) {
|
|
||||||
cases := []struct {
|
|
||||||
err error
|
|
||||||
want string
|
|
||||||
}{
|
|
||||||
{&backend.RateLimitError{Message: "x"}, "rate_limit"},
|
|
||||||
{&backend.ToolUnsupportedError{Message: "x"}, "tool_unsupported"},
|
|
||||||
{&backend.ContextOverflowError{Message: "x"}, "context_overflow"},
|
|
||||||
{&backend.TransientError{Message: "x"}, "transient"},
|
|
||||||
}
|
|
||||||
for _, tc := range cases {
|
|
||||||
got := classifyError(tc.err)
|
|
||||||
if got != tc.want {
|
|
||||||
t.Errorf("classifyError(%T) = %q; want %q", tc.err, got, tc.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
1178
session/session.go
1178
session/session.go
File diff suppressed because it is too large
Load Diff
|
|
@ -1,168 +0,0 @@
|
||||||
package session
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
"github.com/simonfxr/pubsub"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
olog "ollie/log"
|
|
||||||
)
|
|
||||||
|
|
||||||
// newTestCore returns a minimal Session for testing.
|
|
||||||
func newTestCore(initialState string) *Session {
|
|
||||||
ag := &Agent{
|
|
||||||
state: initialState,
|
|
||||||
}
|
|
||||||
ag.InitCond()
|
|
||||||
a := &Session{
|
|
||||||
id: "test",
|
|
||||||
bus: pubsub.NewBus(),
|
|
||||||
env: make(map[string]string),
|
|
||||||
r: ag,
|
|
||||||
log: olog.NewWriter("test", olog.LevelError+1, io.Discard, io.Discard),
|
|
||||||
}
|
|
||||||
return a
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWaitChange_ReturnOnChange verifies that WaitChange unblocks when
|
|
||||||
// setState is called with a different value.
|
|
||||||
func TestWaitChange_ReturnOnChange(t *testing.T) {
|
|
||||||
a := newTestCore("idle")
|
|
||||||
|
|
||||||
result := make(chan string, 1)
|
|
||||||
go func() {
|
|
||||||
v, ok := a.WaitChange(context.Background(), WatchState, "idle")
|
|
||||||
if !ok {
|
|
||||||
result <- "!ok"
|
|
||||||
return
|
|
||||||
}
|
|
||||||
result <- v
|
|
||||||
}()
|
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond) // let goroutine reach cond.Wait
|
|
||||||
a.setState("thinking")
|
|
||||||
|
|
||||||
select {
|
|
||||||
case got := <-result:
|
|
||||||
if got != "thinking" {
|
|
||||||
t.Errorf("WaitChange returned %q; want %q", got, "thinking")
|
|
||||||
}
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("WaitChange did not unblock after setState")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWaitChange_AlreadyChanged verifies that if the value has already
|
|
||||||
// changed before WaitChange is called, it returns immediately.
|
|
||||||
func TestWaitChange_AlreadyChanged(t *testing.T) {
|
|
||||||
a := newTestCore("thinking")
|
|
||||||
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() {
|
|
||||||
v, ok := a.WaitChange(context.Background(), WatchState, "idle")
|
|
||||||
if !ok || v != "thinking" {
|
|
||||||
t.Errorf("WaitChange returned (%q, %v); want (\"thinking\", true)", v, ok)
|
|
||||||
}
|
|
||||||
close(done)
|
|
||||||
}()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("WaitChange blocked when value already changed")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWaitChange_ContextCancel verifies that WaitChange returns ("", false)
|
|
||||||
// when the context is cancelled.
|
|
||||||
func TestWaitChange_ContextCancel(t *testing.T) {
|
|
||||||
a := newTestCore("idle")
|
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
result := make(chan bool, 1)
|
|
||||||
go func() {
|
|
||||||
_, ok := a.WaitChange(ctx, WatchState, "idle")
|
|
||||||
result <- ok
|
|
||||||
}()
|
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
|
||||||
cancel()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case ok := <-result:
|
|
||||||
if ok {
|
|
||||||
t.Error("WaitChange returned ok=true after context cancel; want false")
|
|
||||||
}
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("WaitChange did not unblock after context cancel")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWaitChange_FullCycle simulates idle→thinking→idle and checks each
|
|
||||||
// transition is observed in order.
|
|
||||||
func TestWaitChange_FullCycle(t *testing.T) {
|
|
||||||
a := newTestCore("idle")
|
|
||||||
|
|
||||||
// Step 1: wait for idle→thinking
|
|
||||||
thinking := make(chan string, 1)
|
|
||||||
go func() {
|
|
||||||
v, _ := a.WaitChange(context.Background(), WatchState, "idle")
|
|
||||||
thinking <- v
|
|
||||||
}()
|
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
|
||||||
a.setState("thinking")
|
|
||||||
|
|
||||||
var got string
|
|
||||||
select {
|
|
||||||
case got = <-thinking:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("did not observe idle→thinking")
|
|
||||||
}
|
|
||||||
if got != "thinking" {
|
|
||||||
t.Errorf("step1: got %q; want \"thinking\"", got)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Step 2: wait for thinking→idle
|
|
||||||
idle := make(chan string, 1)
|
|
||||||
go func() {
|
|
||||||
v, _ := a.WaitChange(context.Background(), WatchState, "thinking")
|
|
||||||
idle <- v
|
|
||||||
}()
|
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
|
||||||
a.setState("idle")
|
|
||||||
|
|
||||||
select {
|
|
||||||
case got = <-idle:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("did not observe thinking→idle")
|
|
||||||
}
|
|
||||||
if got != "idle" {
|
|
||||||
t.Errorf("step2: got %q; want \"idle\"", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestWaitChange_NoMissedWakeup fires setState concurrently with WaitChange
|
|
||||||
// to stress the missed-wakeup scenario.
|
|
||||||
func TestWaitChange_NoMissedWakeup(t *testing.T) {
|
|
||||||
const rounds = 500
|
|
||||||
for i := 0; i < rounds; i++ {
|
|
||||||
a := newTestCore("idle")
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() {
|
|
||||||
a.WaitChange(context.Background(), WatchState, "idle") //nolint:errcheck
|
|
||||||
close(done)
|
|
||||||
}()
|
|
||||||
// setState races with WaitChange entering the wait loop.
|
|
||||||
a.setState("thinking")
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatalf("round %d: WaitChange missed wakeup", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Reference in New Issue