From 8e4cdaa4bd25648cb151149f2b2b599f6bdf3308 Mon Sep 17 00:00:00 2001 From: Levi Neely Date: Wed, 29 Jul 2026 21:43:49 +0200 Subject: [PATCH] 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. --- agent/agent.go | 511 +++++ {session => agent}/agent_config.go | 2 +- agent/build_runtime.go | 221 +++ agent/commands.go | 364 ++++ {session => agent}/compaction.go | 2 +- agent/config_paths.go | 41 + {session => agent}/cost.go | 2 +- {session => agent}/fifo.go | 2 +- {session => agent}/fifo_test.go | 2 +- {session => agent}/history.go | 36 +- {session => agent}/hooks.go | 2 +- {session => agent}/loop.go | 9 +- agent/new.go | 60 + {session => agent}/prompt_resolver.go | 2 +- {session => agent}/runtime.go | 2 +- {session => agent}/state.go | 2 +- {session => agent}/turn.go | 28 +- {session => agent}/usage_log.go | 2 +- session/agent.go | 210 -- session/commands.go | 336 +--- session/compaction_test.go | 113 -- session/config_test.go | 137 -- session/core_test.go | 2595 ------------------------- session/cost_test.go | 100 - session/loop_cache_test.go | 187 -- session/loop_errlimit_test.go | 178 -- session/loop_maxsteps_helpers_test.go | 76 - session/loop_maxsteps_test.go | 184 -- session/loop_parallel_test.go | 206 -- session/loop_truncate_test.go | 126 -- session/loop_turnerror_test.go | 164 -- session/session.go | 1178 +++-------- session/waitchange_test.go | 168 -- 33 files changed, 1539 insertions(+), 5709 deletions(-) create mode 100644 agent/agent.go rename {session => agent}/agent_config.go (99%) create mode 100644 agent/build_runtime.go create mode 100644 agent/commands.go rename {session => agent}/compaction.go (98%) create mode 100644 agent/config_paths.go rename {session => agent}/cost.go (99%) rename {session => agent}/fifo.go (97%) rename {session => agent}/fifo_test.go (98%) rename {session => agent}/history.go (94%) rename {session => agent}/hooks.go (99%) rename {session => agent}/loop.go (99%) create mode 100644 agent/new.go rename {session => agent}/prompt_resolver.go (99%) rename {session => agent}/runtime.go (98%) rename {session => agent}/state.go (99%) rename {session => agent}/turn.go (94%) rename {session => agent}/usage_log.go (99%) delete mode 100644 session/agent.go delete mode 100644 session/compaction_test.go delete mode 100644 session/config_test.go delete mode 100644 session/core_test.go delete mode 100644 session/cost_test.go delete mode 100644 session/loop_cache_test.go delete mode 100644 session/loop_errlimit_test.go delete mode 100644 session/loop_maxsteps_helpers_test.go delete mode 100644 session/loop_maxsteps_test.go delete mode 100644 session/loop_parallel_test.go delete mode 100644 session/loop_truncate_test.go delete mode 100644 session/loop_turnerror_test.go delete mode 100644 session/waitchange_test.go diff --git a/agent/agent.go b/agent/agent.go new file mode 100644 index 0000000..e2c2206 --- /dev/null +++ b/agent/agent.go @@ -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: " + 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() +} diff --git a/session/agent_config.go b/agent/agent_config.go similarity index 99% rename from session/agent_config.go rename to agent/agent_config.go index fc8f1db..30e7442 100644 --- a/session/agent_config.go +++ b/agent/agent_config.go @@ -1,4 +1,4 @@ -package session +package agent import ( "encoding/json" diff --git a/agent/build_runtime.go b/agent/build_runtime.go new file mode 100644 index 0000000..35024b9 --- /dev/null +++ b/agent/build_runtime.go @@ -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 +} diff --git a/agent/commands.go b/agent/commands.go new file mode 100644 index 0000000..cc7f289 --- /dev/null +++ b/agent/commands.go @@ -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 +} diff --git a/session/compaction.go b/agent/compaction.go similarity index 98% rename from session/compaction.go rename to agent/compaction.go index 836e251..06a8654 100644 --- a/session/compaction.go +++ b/agent/compaction.go @@ -1,4 +1,4 @@ -package session +package agent import "os" diff --git a/agent/config_paths.go b/agent/config_paths.go new file mode 100644 index 0000000..7e75a19 --- /dev/null +++ b/agent/config_paths.go @@ -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" +} diff --git a/session/cost.go b/agent/cost.go similarity index 99% rename from session/cost.go rename to agent/cost.go index 301e9d0..d0686ac 100644 --- a/session/cost.go +++ b/agent/cost.go @@ -1,4 +1,4 @@ -package session +package agent import ( "strings" diff --git a/session/fifo.go b/agent/fifo.go similarity index 97% rename from session/fifo.go rename to agent/fifo.go index 8a8f566..a4f5dfb 100644 --- a/session/fifo.go +++ b/agent/fifo.go @@ -1,4 +1,4 @@ -package session +package agent import "sync" diff --git a/session/fifo_test.go b/agent/fifo_test.go similarity index 98% rename from session/fifo_test.go rename to agent/fifo_test.go index e9d5fdb..4bb0328 100644 --- a/session/fifo_test.go +++ b/agent/fifo_test.go @@ -1,4 +1,4 @@ -package session +package agent import ( "sync" diff --git a/session/history.go b/agent/history.go similarity index 94% rename from session/history.go rename to agent/history.go index 982d0a0..a039b82 100644 --- a/session/history.go +++ b/agent/history.go @@ -1,13 +1,15 @@ -package session +package agent import ( "context" "encoding/json" + "crypto/rand" "fmt" "os" "slices" "strings" "time" + "strconv" "ollie/backend" ) @@ -694,3 +696,35 @@ func (s *History) contextDebug() 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 + } +} diff --git a/session/hooks.go b/agent/hooks.go similarity index 99% rename from session/hooks.go rename to agent/hooks.go index 45771a9..26e374a 100644 --- a/session/hooks.go +++ b/agent/hooks.go @@ -1,4 +1,4 @@ -package session +package agent import ( "bytes" diff --git a/session/loop.go b/agent/loop.go similarity index 99% rename from session/loop.go rename to agent/loop.go index 06610be..ce2242f 100644 --- a/session/loop.go +++ b/agent/loop.go @@ -1,4 +1,4 @@ -package session +package agent import ( "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 diff --git a/agent/new.go b/agent/new.go new file mode 100644 index 0000000..9533447 --- /dev/null +++ b/agent/new.go @@ -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 +} diff --git a/session/prompt_resolver.go b/agent/prompt_resolver.go similarity index 99% rename from session/prompt_resolver.go rename to agent/prompt_resolver.go index e851adf..dab0643 100644 --- a/session/prompt_resolver.go +++ b/agent/prompt_resolver.go @@ -1,4 +1,4 @@ -package session +package agent import ( "bytes" diff --git a/session/runtime.go b/agent/runtime.go similarity index 98% rename from session/runtime.go rename to agent/runtime.go index 6c83c4f..5074a05 100644 --- a/session/runtime.go +++ b/agent/runtime.go @@ -1,4 +1,4 @@ -package session +package agent import ( "encoding/json" diff --git a/session/state.go b/agent/state.go similarity index 99% rename from session/state.go rename to agent/state.go index 92623f1..3e9ce93 100644 --- a/session/state.go +++ b/agent/state.go @@ -1,4 +1,4 @@ -package session +package agent import ( "context" diff --git a/session/turn.go b/agent/turn.go similarity index 94% rename from session/turn.go rename to agent/turn.go index 3da5ee9..464f42b 100644 --- a/session/turn.go +++ b/agent/turn.go @@ -1,4 +1,4 @@ -package session +package agent import ( "context" @@ -464,3 +464,29 @@ func (ag *Agent) runCompact(ctx context.Context, trigger string) (int, error) { } 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 +} diff --git a/session/usage_log.go b/agent/usage_log.go similarity index 99% rename from session/usage_log.go rename to agent/usage_log.go index 9f262b4..b9e9cbc 100644 --- a/session/usage_log.go +++ b/agent/usage_log.go @@ -1,4 +1,4 @@ -package session +package agent import ( "encoding/json" diff --git a/session/agent.go b/session/agent.go deleted file mode 100644 index a06eabb..0000000 --- a/session/agent.go +++ /dev/null @@ -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: " - 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 -} diff --git a/session/commands.go b/session/commands.go index 9dcaa2d..ff473ec 100644 --- a/session/commands.go +++ b/session/commands.go @@ -6,10 +6,9 @@ import ( "fmt" "os" "path/filepath" - "slices" - "strconv" "strings" + "ollie/agent" ) 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] 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) { if h := a.listHandlers[name]; h != nil { for _, item := range h() { - a.emit(infoEvent(" " + item)) + a.emit(agent.Event{Role: "info", Content: " " + item + "\n"}) } } } type cmdFn func([]string) 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) { allFlag := len(args) > 0 && args[0] == "-a" type sessionFile struct { @@ -343,7 +71,7 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool { if readErr != nil { continue } - var ps PersistedAgent + var ps agent.PersistedAgent if json.Unmarshal(data, &ps) != nil { continue } @@ -364,77 +92,73 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool { if len(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 } 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) { - if a.r.history == nil { - a.emit(infoEvent("error: no active session")) + if !a.r.HasHistory() { + a.emit(agent.Event{Role: "info", Content: "error: no active session\n"}) return } 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 } name := args[0] path := a.sessionsDir + "/" + name + ".json" - if err := a.r.history.saveTo(path, name, a.r.agentName, a.CWD()); err != nil { - a.emit(infoEvent("error: " + err.Error())) + if err := a.r.SaveTo(path, name, a.CWD()); err != nil { + a.emit(agent.Event{Role: "info", Content: "error: " + err.Error() + "\n"}) return } - a.emit(infoEvent("saved: " + path)) + a.emit(agent.Event{Role: "info", Content: "saved: " + path + "\n"}) }, "/resume": func(args []string) { 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 } 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 } name := args[0] path := a.sessionsDir + "/" + name + ".json" data, err := os.ReadFile(path) 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 } - var ps PersistedAgent + var ps agent.PersistedAgent 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 } - a.r.history = RestoreHistory(&ps) - a.emit(infoEvent(fmt.Sprintf("resumed session %s (%d messages)", name, len(ps.Messages)))) + a.r.Restore(&ps) + a.emit(agent.Event{Role: "info", Content: fmt.Sprintf("resumed session %s (%d messages)\n", name, len(ps.Messages))}) }, "/cwd": func(args []string) { if len(args) == 0 { - a.emit(infoEvent("cwd: " + a.CWD())) + a.emit(agent.Event{Role: "info", Content: "cwd: " + a.CWD() + "\n"}) return } dir := strings.Join(args, " ") 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 } - a.emit(infoEvent("cwd: " + dir)) + a.emit(agent.Event{Role: "info", Content: "cwd: " + dir + "\n"}) }, "/skills": func(args []string) { listFromHandler("skills") }, "/tools": func(args []string) { listFromHandler("tools") }, - "/sp": func(args []string) { - a.emit(infoEvent(a.r.runtime.Preamble)) - }, - "/help": func(args []string) { lines := []string{ "Available commands:", @@ -452,21 +176,17 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool { " /cwd [path] - show or change working directory", " /i - inject prompt into the running turn", " /irw - rewrite the pending inject", - " /queued [pop|clear] - manage queued prompts", " /compact - summarize conversation and compact context", " /context - show context size and message breakdown", " /cost - show last turn and session cost", " /usage - show token usage and context percentage", " /history - dump bounded message history", " /clear - clear session", - " /kill - kill session", - " /rn - rename session", " /sp - show rendered system prompt", " /help - show this help", - " ! - run shell command", } 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 { return false } - a.emit(infoEvent("")) + a.emit(agent.Event{Role: "info", Content: "\n"}) fn(args) return true } diff --git a/session/compaction_test.go b/session/compaction_test.go deleted file mode 100644 index 70f40ef..0000000 --- a/session/compaction_test.go +++ /dev/null @@ -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) - } - } - } -} diff --git a/session/config_test.go b/session/config_test.go deleted file mode 100644 index 81cea54..0000000 --- a/session/config_test.go +++ /dev/null @@ -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") - } -} diff --git a/session/core_test.go b/session/core_test.go deleted file mode 100644 index 42cd8f3..0000000 --- a/session/core_test.go +++ /dev/null @@ -1,2595 +0,0 @@ -package session - -import ( - "context" - "encoding/json" - "fmt" - "os" - "path/filepath" - "strings" - "sync" - "sync/atomic" - "testing" - "time" - - "ollie/backend" - "ollie/tools" -) - -// --- mock backend --- - -type mockBackend struct { - mu sync.Mutex - name string - model string - ctxLen int - models []string - respond func(context.Context, []backend.Message, []backend.Tool, backend.GenerationParams) (<-chan backend.StreamEvent, error) -} - -func (m *mockBackend) ChatStream(ctx context.Context, msgs []backend.Message, ts []backend.Tool, p backend.GenerationParams) (<-chan backend.StreamEvent, error) { - m.mu.Lock() - fn := m.respond - m.mu.Unlock() - if fn != nil { - return fn(ctx, msgs, ts, p) - } - return textStream("ok"), nil -} -func (m *mockBackend) Name() string { return m.name } -func (m *mockBackend) DefaultModel() string { return m.model } -func (m *mockBackend) Model() string { m.mu.Lock(); defer m.mu.Unlock(); return m.model } -func (m *mockBackend) SetModel(s string) { m.mu.Lock(); m.model = s; m.mu.Unlock() } -func (m *mockBackend) ContextLength(_ context.Context) int { return m.ctxLen } -func (m *mockBackend) Models(_ context.Context) []string { - m.mu.Lock() - defer m.mu.Unlock() - return m.models -} - -// textStream returns a single-event stream carrying the given text. -func textStream(text string) <-chan backend.StreamEvent { - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{ - Content: text, - Done: true, - StopReason: "stop", - Usage: backend.Usage{InputTokens: 10, OutputTokens: 5}, - } - close(ch) - return ch -} - -// blockedStream returns a channel that delivers a response when unblock is -// closed, or drains silently if ctx is cancelled first. -func blockedStream(ctx context.Context, unblock <-chan struct{}) <-chan backend.StreamEvent { - ch := make(chan backend.StreamEvent, 1) - go func() { - defer close(ch) - select { - case <-ctx.Done(): - case <-unblock: - ch <- backend.StreamEvent{Content: "ok", Done: true, StopReason: "stop"} - } - }() - return ch -} - -// --- test helpers --- - -func defaultBE() *mockBackend { - return &mockBackend{name: "mock", model: "test", ctxLen: 128000} -} - -// newCore builds a minimal *agent for tests, bypassing loadSystemPrompt -// by directly setting Preamble on Runtime. -func newCore(t *testing.T, be backend.Backend, hooks Hooks) *Session { - t.Helper() - t.Setenv("OLLIE", "") - if be == nil { - be = defaultBE() - } - if hooks == nil { - hooks = Hooks{} - } - env := &Runtime{ - Hooks: hooks, - Preamble: "test system prompt", - } - c := New(Config{ - Backend: be, - AgentName: "test", - AgentsDir: t.TempDir(), - SessionsDir: t.TempDir(), - SessionID: NewSessionID(), - CWD: t.TempDir(), - Runtime: env, - NewDispatcher: tools.NewDispatcher, - }) - t.Cleanup(c.Close) - return c -} - -// collectEvents runs Submit synchronously and returns all emitted events. -func collectEvents(ctx context.Context, c *Session, input string) []Event { - var mu sync.Mutex - var evs []Event - sub := c.Bus().Subscribe("event", func(ev Event) { - mu.Lock() - evs = append(evs, ev) - mu.Unlock() - }) - c.Submit(ctx, input) - c.Bus().Unsubscribe(sub) - mu.Lock() - defer mu.Unlock() - return evs -} - -// byRole returns the Content of every event with the given role. -func byRole(evs []Event, role string) []string { - var out []string - for _, ev := range evs { - if ev.Role == role { - out = append(out, ev.Content) - } - } - return out -} - -// waitState blocks until c.State() == want, failing after 2 s. -func waitState(t *testing.T, c *Session, want string) { - t.Helper() - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) - defer cancel() - for c.State() != want { - if _, ok := c.WaitChange(ctx, WatchState, c.State()); !ok { - t.Fatalf("timed out waiting for state %q (current: %q)", want, c.State()) - } - } -} - -// --- Submit: happy path --- - -func TestSubmit_HappyPath(t *testing.T) { - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return textStream("hello back"), nil - } - c := newCore(t, be, nil) - - evs := collectEvents(context.Background(), c, "hello") - - if got := byRole(evs, "user"); len(got) == 0 || got[0] != "hello" { - t.Errorf("user event: %v", got) - } - if got := byRole(evs, "assistant"); len(got) == 0 || got[0] != "hello back" { - t.Errorf("assistant event: %v", got) - } - if got := c.State(); got != "idle" { - t.Errorf("State() = %q; want idle", got) - } - if got := c.Reply(); got != "hello back" { - t.Errorf("Reply() = %q; want %q", got, "hello back") - } - if c.IsRunning() { - t.Error("IsRunning() = true after turn; want false") - } -} - -// --- State transitions --- - -func TestSubmit_StateTransitions(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - if got := c.State(); got != "idle" { - t.Fatalf("initial state = %q; want idle", got) - } - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - - waitState(t, c, "thinking") - close(unblock) - <-done - - if got := c.State(); got != "idle" { - t.Errorf("final state = %q; want idle", got) - } -} - -func TestSubmit_ToolCallStateTransitions(t *testing.T) { - const toolName = "my_tool" - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - if callCount == 1 { - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{ - ToolCalls: []backend.ToolCall{{Name: toolName, Arguments: json.RawMessage(`{}`)}}, - Done: true, - StopReason: "tool_calls", - } - close(ch) - return ch, nil - - } - return textStream("done"), nil - - } - - c := newCore(t, be, nil) - - var stateAtExec string - c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) { - stateAtExec = c.State() - return `{}`, nil, nil - } - - collectEvents(context.Background(), c, "run tool") - - wantState := "calling: " + toolName - if stateAtExec != wantState { - t.Errorf("state during tool exec = %q; want %q", stateAtExec, wantState) - } - if got := c.State(); got != "idle" { - t.Errorf("final state = %q; want idle", got) - } -} - -// --- preTurn hook --- - -func TestSubmit_PreTurnHookBlocks(t *testing.T) { - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - return textStream("should not be called"), nil - - } - c := newCore(t, be, Hooks{HookPreTurn: []string{"exit 2"}}) - - evs := collectEvents(context.Background(), c, "hello") - - if callCount > 0 { - t.Error("backend called despite preTurn hook blocking") - } - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "hook blocked prompt") { - found = true - } - } - if !found { - t.Errorf("expected 'hook blocked prompt' info event; got: %v", byRole(evs, "info")) - } -} - -func TestSubmit_PreTurnHookContext(t *testing.T) { - var lastUserMsg string - be := defaultBE() - be.respond = func(_ context.Context, msgs []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - for _, m := range msgs { - if m.Role == "user" { - lastUserMsg = m.Content - } - } - return textStream("ok"), nil - } - c := newCore(t, be, Hooks{HookPreTurn: []string{`echo "extra context"`}}) - - collectEvents(context.Background(), c, "base prompt") - - if !strings.Contains(lastUserMsg, "extra context") { - t.Errorf("user message %q does not contain hook-injected context", lastUserMsg) - } -} - -// --- postTurn hook continuation --- - -func TestSubmit_PostTurnHookContinue(t *testing.T) { - flagFile := filepath.Join(t.TempDir(), "fired") - hookScript := fmt.Sprintf( - `if [ ! -f %q ]; then touch %q; printf "auto-continue" >&2; exit 2; fi`, - flagFile, flagFile, - ) - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - return textStream(fmt.Sprintf("response %d", callCount)), nil - - } - c := newCore(t, be, Hooks{HookPostTurn: []string{hookScript}}) - - collectEvents(context.Background(), c, "first prompt") - - if callCount != 2 { - t.Errorf("backend called %d times; want 2 (original + hook continuation)", callCount) - } -} - -// --- FIFO drain --- - -func TestSubmit_FIFODrain(t *testing.T) { - var callCount atomic.Int32 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount.Add(1) - return textStream("ok"), nil - } - c := newCore(t, be, nil) - - c.Queue("second") - c.Queue("third") - c.Submit(context.Background(), "first") - - // Wait for async drain goroutines to complete. - deadline := time.Now().Add(2 * time.Second) - for callCount.Load() < 3 && time.Now().Before(deadline) { - time.Sleep(5 * time.Millisecond) - } - - if got := callCount.Load(); got != 3 { - t.Errorf("backend called %d times; want 3 (first + second + third)", got) - } -} - -// --- pendingInject as next turn --- - -// TestSubmit_PendingInjectAsNextTurn verifies that a pendingInject left -// unconsumed (no tool calls) becomes the next prompt after the turn. -func TestSubmit_PendingInjectAsNextTurn(t *testing.T) { - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - return textStream("ok"), nil - } - c := newCore(t, be, nil) - - inject := "injected follow-up" - c.r.pendingInject.Store(&inject) - - collectEvents(context.Background(), c, "first") - - if callCount != 2 { - t.Errorf("backend called %d times; want 2 (first turn + inject turn)", callCount) - } -} - -// --- Interrupt --- - -func TestInterrupt_Running(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - - waitState(t, c, "thinking") - - if !c.Interrupt(ErrInterrupted) { - t.Error("Interrupt() = false; want true (action was running)") - } - - select { - case <-done: - case <-time.After(2 * time.Second): - t.Fatal("Submit did not return after Interrupt") - } - if got := c.State(); got != "idle" { - t.Errorf("State() = %q after interrupt; want idle", got) - } - if c.IsRunning() { - t.Error("IsRunning() = true after interrupt; want false") - } -} - -func TestInterrupt_Idle(t *testing.T) { - c := newCore(t, nil, nil) - if c.Interrupt(ErrInterrupted) { - t.Error("Interrupt() = true when idle; want false") - } -} - -// --- Submit while running --- - -func TestSubmit_WhileRunning(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "first") - }() - - waitState(t, c, "thinking") - - // Concurrent Submit must not start a new turn — always goes to FIFO. - c.Submit(context.Background(), "concurrent") - - if stored := c.r.pendingInject.Load(); stored != nil { - t.Errorf("concurrent Submit set pendingInject %q; want FIFO only", *stored) - } - got, inFIFO := c.PopQueue() - if !inFIFO || got != "concurrent" { - t.Errorf("concurrent Submit PopQueue = %q, %v; want %q, true", got, inFIFO, "concurrent") - } - - close(unblock) - <-done -} - -// --- compaction --- - -func TestManualCompact(t *testing.T) { - callCount := 0 - var stateAtCompact string - be := defaultBE() - c := newCore(t, be, nil) - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - stateAtCompact = c.State() - return textStream("summary or answer"), nil - - } - - // Seed a session with more than hotTailSize+warmIndexSize messages so compact() doesn't short-circuit. - c.r.history = newHistory("goal") - for i := range 15 { - c.r.history.messages = append(c.r.history.messages, - backend.Message{Role: "assistant", Content: fmt.Sprintf("response %d with enough text to count", i)}, - backend.Message{Role: "user", Content: fmt.Sprintf("follow up %d", i)}, - ) - } - before := len(c.r.history.messages) - - evs := collectEvents(context.Background(), c, "/compact") - - if c.r.history == nil { - t.Fatal("session nil after compact") - } - if after := len(c.r.history.messages); after >= before { - t.Errorf("messages: before=%d after=%d; want fewer after compact", before, after) - } - if callCount == 0 { - t.Error("backend not called for compaction") - } - if stateAtCompact != "compacting" { - t.Errorf("state during compact = %q; want compacting", stateAtCompact) - } - if got := c.State(); got != "idle" { - t.Errorf("state after compact = %q; want idle", got) - } - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "compacted") { - found = true - } - } - if !found { - t.Errorf("no 'compacted' info event; got: %v", byRole(evs, "info")) - } -} - -func TestAutoCompact(t *testing.T) { - callCount := 0 - var stateAtCompact string - // ctxLen=10 → autoCompactLimit = 7 tokens; our seeded session exceeds this. - be := &mockBackend{name: "mock", model: "test", ctxLen: 10} - c := newCore(t, be, nil) - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - if callCount == 1 { - stateAtCompact = c.State() - return textStream("summary text for compaction"), nil - - } - return textStream("answer"), nil - - } - - c.r.history = newHistory("goal") - for range 15 { - c.r.history.messages = append(c.r.history.messages, - backend.Message{Role: "assistant", Content: "a long response that exceeds seven tokens of content"}, - backend.Message{Role: "user", Content: "follow up question with enough text to push over the limit"}, - ) - } - - collectEvents(context.Background(), c, "next prompt") - - if callCount < 2 { - t.Errorf("backend called %d times; want ≥2 (compact + turn)", callCount) - } - if stateAtCompact != "compacting" { - t.Errorf("state during auto-compact = %q; want compacting", stateAtCompact) - } -} - -// --- agentSpawn --- - -func TestAgentSpawnFiresOnce(t *testing.T) { - logFile := filepath.Join(t.TempDir(), "spawned") - hookScript := fmt.Sprintf(`printf "spawn\n" >> %q`, logFile) - c := newCore(t, nil, Hooks{HookAgentSpawn: []string{hookScript}}) - - collectEvents(context.Background(), c, "first") - collectEvents(context.Background(), c, "second") - - data, err := os.ReadFile(logFile) - if err != nil { - t.Fatalf("spawn log not written: %v", err) - } - lines := strings.Count(string(data), "\n") - if lines != 1 { - t.Errorf("agentSpawn hook fired %d time(s); want 1", lines) - } -} - -// --- panic recovery --- - -func TestSubmit_PanicRecovery(t *testing.T) { - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - panic("backend exploded") - } - c := newCore(t, be, nil) - - var errEvents []Event - sub := c.Bus().Subscribe("event", func(ev Event) { - if ev.Role == "error" { - errEvents = append(errEvents, ev) - } - }) - c.Submit(context.Background(), "hello") - c.Bus().Unsubscribe(sub) - - if len(errEvents) == 0 { - t.Error("no error event after panic; expected one") - } - if got := c.State(); got != "idle" { - t.Errorf("State() = %q after panic; want idle", got) - } - if c.IsRunning() { - t.Error("IsRunning() = true after panic; want false") - } -} - -// --- commands --- - -func TestCommand_Help(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/help") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "Available commands") { - found = true - } - } - if !found { - t.Errorf("/help: 'Available commands' not found in info events: %v", byRole(evs, "info")) - } -} - -func TestCommand_Clear(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "first turn") - if c.r.history == nil { - t.Fatal("session nil after first turn") - } - oldID := c.id - collectEvents(context.Background(), c, "/clear") - if c.r.history != nil { - t.Error("session not nil after /clear") - } - if c.id != oldID { - t.Errorf("sessionID changed after /clear: %q → %q; want unchanged", oldID, c.id) - } -} - -// --- CWD --- - -func TestSetCWD(t *testing.T) { - c := newCore(t, nil, nil) - dir := t.TempDir() - if err := c.SetCWD(dir); err != nil { - t.Fatalf("SetCWD(%q): %v", dir, err) - } - if got := c.CWD(); got != dir { - t.Errorf("CWD() = %q; want %q", got, dir) - } -} - -func TestSetCWD_NonExistent(t *testing.T) { - c := newCore(t, nil, nil) - if err := c.SetCWD("/nonexistent/path/xyz/abc"); err == nil { - t.Error("SetCWD with nonexistent path returned nil; want error") - } -} - -// --- Queue / PopQueue --- - -func TestQueuePopQueue(t *testing.T) { - c := newCore(t, nil, nil) - c.Queue("a") - c.Queue("b") - - if got, ok := c.PopQueue(); !ok || got != "a" { - t.Errorf("PopQueue() = %q, %v; want %q, true", got, ok, "a") - } - if got, ok := c.PopQueue(); !ok || got != "b" { - t.Errorf("PopQueue() = %q, %v; want %q, true", got, ok, "b") - } - if got, ok := c.PopQueue(); ok { - t.Errorf("PopQueue on empty = %q, true; want empty, false", got) - } -} - -// --- /i and /irw --- - -func TestCommand_I_Empty(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/i") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/i with no args: expected error event, got: %v", byRole(evs, "info")) - } -} - -func TestCommand_I_SetsInject(t *testing.T) { - // When running, /i sets pendingInject. - be := defaultBE() - done := make(chan struct{}) - be.respond = func(_ context.Context, msgs []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - <-done - return textStream("ok"), nil - } - c := newCore(t, be, nil) - submitDone := make(chan struct{}) - go func() { - c.Submit(context.Background(), "hello") - close(submitDone) - }() - waitState(t, c, "thinking") - c.handleCommand(context.Background(), "/i my inject") - p := c.r.pendingInject.Load() - if p == nil || *p != "my inject" { - t.Errorf("pendingInject = %v; want %q", p, "my inject") - } - close(done) - <-submitDone -} - -func TestCommand_I_SubmitsWhenIdle(t *testing.T) { - // When idle, /i submits the prompt directly. - var callCount atomic.Int32 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount.Add(1) - return textStream("ok"), nil - } - c := newCore(t, be, nil) - c.Submit(context.Background(), "/i my prompt") - deadline := time.Now().Add(2 * time.Second) - for callCount.Load() < 1 && time.Now().Before(deadline) { - time.Sleep(5 * time.Millisecond) - } - if callCount.Load() != 1 { - t.Errorf("expected 1 backend call from /i submit; got %d", callCount.Load()) - } -} - -func TestCommand_IRW_Empty(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/irw") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/irw with no args: expected error event, got: %v", byRole(evs, "info")) - } -} - -func TestCommand_IRW_OverwritesInject(t *testing.T) { - c := newCore(t, nil, nil) - existing := "old" - c.r.pendingInject.Store(&existing) - collectEvents(context.Background(), c, "/irw new inject") - p := c.r.pendingInject.Load() - if p == nil || *p != "new inject" { - t.Errorf("pendingInject after /irw = %v; want %q", p, "new inject") - } -} - -// --- /compact additional paths --- - -func TestCommand_Compact_WhileRunning(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - waitState(t, c, "thinking") - - evs := collectEvents(context.Background(), c, "/compact") - - close(unblock) - <-done - - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/compact while running: expected error event, got: %v", byRole(evs, "info")) - } -} - -func TestCommand_Compact_NilSession(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/compact") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "nothing to compact") { - found = true - } - } - if !found { - t.Errorf("/compact with nil session: expected 'nothing to compact', got: %v", byRole(evs, "info")) - } -} - -func TestCommand_Compact_PreHookBlocks(t *testing.T) { - c := newCore(t, nil, Hooks{HookPreCompact: []string{"exit 2"}}) - c.r.history = newHistory("goal") - for i := range 5 { - c.r.history.messages = append(c.r.history.messages, - backend.Message{Role: "assistant", Content: fmt.Sprintf("response %d", i)}, - backend.Message{Role: "user", Content: fmt.Sprintf("follow up %d", i)}, - ) - } - evs := collectEvents(context.Background(), c, "/compact") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "cancelled") { - found = true - } - } - if !found { - t.Errorf("/compact with blocking preHook: expected 'cancelled', got: %v", byRole(evs, "info")) - } -} - -// --- /clear while running --- - -func TestCommand_Clear_WhileRunning(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - waitState(t, c, "thinking") - - evs := collectEvents(context.Background(), c, "/clear") - - close(unblock) - <-done - - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/clear while running: expected error event, got: %v", byRole(evs, "info")) - } -} - -// --- executeTurn: backend error --- - -func TestSubmit_BackendError(t *testing.T) { - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return nil, fmt.Errorf("backend unavailable") - } - c := newCore(t, be, nil) - evs := collectEvents(context.Background(), c, "hello") - found := false - for _, ev := range evs { - if ev.Role == "error" && strings.Contains(ev.Content, "backend unavailable") { - found = true - } - } - if !found { - t.Errorf("backend error: expected error event, got: %v", evs) - } - if got := c.State(); got != "idle" { - t.Errorf("State() = %q after backend error; want idle", got) - } -} - -// --- executeTurn: startup messages --- - -func TestSubmit_StartupMessages(t *testing.T) { - c := newCore(t, nil, nil) - c.r.startupMessages = []string{"startup msg 1", "startup msg 2"} - evs := collectEvents(context.Background(), c, "hello") - found := 0 - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "startup msg") { - found++ - } - } - if found != 2 { - t.Errorf("startup messages: found %d info events; want 2; got: %v", found, byRole(evs, "info")) - } - if c.r.startupMessages != nil { - t.Error("startupMessages not cleared after first turn") - } -} - -// --- loop: stream interrupted without Done --- - -func TestRun_StreamInterrupted(t *testing.T) { - old := streamDropBaseDelay - streamDropBaseDelay = 10 * time.Millisecond - defer func() { streamDropBaseDelay = old }() - - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - ch := make(chan backend.StreamEvent) - close(ch) // close without sending Done=true - return ch, nil - - } - c := newCore(t, be, nil) - evs := collectEvents(context.Background(), c, "hello") - found := false - for _, ev := range evs { - if ev.Role == "error" { - found = true - } - } - if !found { - t.Errorf("stream interrupted: expected error event, got: %v", evs) - } -} - -// --- loop: unknown stop reason --- - -func TestRun_UnknownStopReason(t *testing.T) { - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{Done: true, StopReason: "max_completion_tokens", Content: "partial"} - close(ch) - return ch, nil - - } - c := newCore(t, be, nil) - evs := collectEvents(context.Background(), c, "hello") - found := false - for _, ev := range evs { - if ev.Role == "error" { - found = true - } - } - if !found { - t.Errorf("unknown stop reason: expected error event, got: %v", evs) - } -} - -// --- loop: tool call with empty name --- - -func TestRun_ToolEmptyName(t *testing.T) { - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, msgs []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - if callCount == 1 { - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{ - ToolCalls: []backend.ToolCall{{Name: "", Arguments: json.RawMessage(`{}`)}}, - Done: true, - StopReason: "tool_calls", - } - close(ch) - return ch, nil - - } - return textStream("done"), nil - - } - c := newCore(t, be, nil) - collectEvents(context.Background(), c, "run empty tool") - if callCount != 2 { - t.Errorf("backend called %d times; want 2 (empty-tool + follow-up)", callCount) - } - if got := c.State(); got != "idle" { - t.Errorf("State() = %q after empty-tool turn; want idle", got) - } -} - -// --- loop: no tool executor configured --- - -func TestRun_NoExec(t *testing.T) { - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - if callCount == 1 { - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{ - ToolCalls: []backend.ToolCall{{Name: "my_tool", Arguments: json.RawMessage(`{}`)}}, - Done: true, - StopReason: "tool_calls", - } - close(ch) - return ch, nil - - } - return textStream("done"), nil - - } - c := newCore(t, be, nil) - c.r.runtime.Exec = nil // no executor - collectEvents(context.Background(), c, "run tool") - if callCount != 2 { - t.Errorf("backend called %d times; want 2", callCount) - } -} - -// --- hooks: non-zero non-two exit code --- - -func TestHook_NonZeroExitCode(t *testing.T) { - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - return textStream("ok"), nil - } - // exit 1 is a non-blocking warning: turn should still proceed. - c := newCore(t, be, Hooks{HookPreTurn: []string{"exit 1"}}) - collectEvents(context.Background(), c, "hello") - if callCount == 0 { - t.Error("backend not called after exit-1 preTurn hook; want call (non-blocking)") - } -} - -// --- hooks: context cancelled mid-hook --- - -func TestHook_ContextCancelled(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) - defer cancel() - // Sleep-10 hook is killed when ctx times out; must not deadlock. - c := newCore(t, nil, Hooks{HookPreTurn: []string{"sleep 10"}}) - collectEvents(ctx, c, "hello") // returns when ctx expires -} - -// --- /backend --- - -func TestCommand_Backend_Show(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/backend") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "mock") { - found = true - } - } - if !found { - t.Errorf("/backend: backend name not in info events: %v", byRole(evs, "info")) - } -} - -func TestCommand_Backend_WhileRunning(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - waitState(t, c, "thinking") - - evs := collectEvents(context.Background(), c, "/backend newbackend") - - close(unblock) - <-done - - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/backend while running: expected error event, got: %v", byRole(evs, "info")) - } -} - -// --- /model --- - -func TestCommand_Model_Show(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/model") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "test") { - found = true - } - } - if !found { - t.Errorf("/model: model name not in info events: %v", byRole(evs, "info")) - } -} - -func TestCommand_Model_Set(t *testing.T) { - be := defaultBE() - c := newCore(t, be, nil) - collectEvents(context.Background(), c, "/model gpt-4") - if got := be.Model(); got != "gpt-4" { - t.Errorf("after /model gpt-4: Model() = %q; want gpt-4", got) - } -} - -func TestCommand_Model_WhileRunning(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - waitState(t, c, "thinking") - - // /model is now allowed while running (needed for turnError hook recovery). - evs := collectEvents(context.Background(), c, "/model new-model") - - close(unblock) - <-done - - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - t.Errorf("/model while running: unexpected error: %s", s) - } - } - if got := be.Model(); got != "new-model" { - t.Errorf("model = %q; want new-model", got) - } -} - -// --- /models --- - -func TestCommand_Models_Empty(t *testing.T) { - c := newCore(t, nil, nil) // defaultBE has nil models - evs := collectEvents(context.Background(), c, "/models") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "no models") { - found = true - } - } - if !found { - t.Errorf("/models with no models: expected 'no models available', got: %v", byRole(evs, "info")) - } -} - -func TestCommand_Models_List(t *testing.T) { - be := &mockBackend{name: "mock", model: "b", ctxLen: 128000, models: []string{"a", "b", "c"}} - c := newCore(t, be, nil) - evs := collectEvents(context.Background(), c, "/models") - infos := byRole(evs, "info") - if len(infos) < 3 { - t.Fatalf("/models: expected ≥3 info events; got: %v", infos) - } - markedCurrent := false - for _, s := range infos { - if strings.Contains(s, "* ") && strings.Contains(s, "b") { - markedCurrent = true - } - } - if !markedCurrent { - t.Errorf("/models: current model 'b' not marked with '* '; got: %v", infos) - } -} - -// --- /agent --- - -func TestCommand_Agent_Show(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/agent") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "test") { - found = true - } - } - if !found { - t.Errorf("/agent: agent name not in info events: %v", byRole(evs, "info")) - } -} - -func TestCommand_Agent_WhileRunning(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - waitState(t, c, "thinking") - - evs := collectEvents(context.Background(), c, "/agent other") - - close(unblock) - <-done - - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/agent while running: expected error event, got: %v", byRole(evs, "info")) - } -} - -func TestCommand_Agent_NotFound(t *testing.T) { - c := newCore(t, nil, nil) // agentsDir is a fresh temp dir - evs := collectEvents(context.Background(), c, "/agent nonexistent") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/agent nonexistent: expected error event, got: %v", byRole(evs, "info")) - } -} - -// --- /agents --- - -func TestCommand_Agents_Empty(t *testing.T) { - c := newCore(t, nil, nil) // agentsDir is a fresh temp dir - t.Setenv("OLLIE_AGENTS_PATH", c.r.agentsDir) - evs := collectEvents(context.Background(), c, "/agents") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "no agents found") { - found = true - } - } - if !found { - t.Errorf("/agents with empty dir: expected 'no agents found', got: %v", byRole(evs, "info")) - } -} - -// --- /sessions --- - -func TestCommand_Sessions_Empty(t *testing.T) { - c := newCore(t, nil, nil) // sessionsDir is a fresh temp dir - evs := collectEvents(context.Background(), c, "/sessions") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "no sessions") { - found = true - } - } - if !found { - t.Errorf("/sessions with empty dir: expected 'no sessions found', got: %v", byRole(evs, "info")) - } -} - -func TestCommand_Sessions_List(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "hello") // creates and saves session - evs := collectEvents(context.Background(), c, "/sessions") - markedCurrent := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "* ") { - markedCurrent = true - } - } - if !markedCurrent { - t.Errorf("/sessions: current session not marked with '* '; got: %v", byRole(evs, "info")) - } -} - -// --- /cwd --- - -func TestCommand_CWD_Show(t *testing.T) { - c := newCore(t, nil, nil) - cwd := c.CWD() - evs := collectEvents(context.Background(), c, "/cwd") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, cwd) { - found = true - } - } - if !found { - t.Errorf("/cwd: path %q not in info events: %v", cwd, byRole(evs, "info")) - } -} - -func TestCommand_CWD_Set(t *testing.T) { - c := newCore(t, nil, nil) - dir := t.TempDir() - evs := collectEvents(context.Background(), c, "/cwd "+dir) - if got := c.CWD(); got != dir { - t.Errorf("CWD() = %q; want %q", got, dir) - } - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, dir) { - found = true - } - } - if !found { - t.Errorf("/cwd set: new path not confirmed in info events: %v", byRole(evs, "info")) - } -} - -func TestCommand_CWD_SetNonExistent(t *testing.T) { - c := newCore(t, nil, nil) - old := c.CWD() - evs := collectEvents(context.Background(), c, "/cwd /nonexistent/xyz/abc") - if got := c.CWD(); got != old { - t.Errorf("CWD changed to %q after invalid path; want unchanged %q", got, old) - } - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "error") { - found = true - } - } - if !found { - t.Errorf("/cwd nonexistent: expected error event, got: %v", byRole(evs, "info")) - } -} - -// --- /context --- - -func TestCommand_Context_NoSession(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/context") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "no active session") { - found = true - } - } - if !found { - t.Errorf("/context before any turn: expected 'no active session', got: %v", byRole(evs, "info")) - } -} - -func TestCommand_Context_WithSession(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "hello") - evs := collectEvents(context.Background(), c, "/context") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "tokens") { - found = true - } - } - if !found { - t.Errorf("/context after turn: expected token usage line, got: %v", byRole(evs, "info")) - } -} - -// --- /usage --- - -func TestCommand_Usage_NoSession(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/usage") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "no active session") { - found = true - } - } - if !found { - t.Errorf("/usage before any turn: expected 'no active session', got: %v", byRole(evs, "info")) - } -} - -func TestCommand_Usage_WithSession(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "hello") - evs := collectEvents(context.Background(), c, "/usage") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "requests") { - found = true - } - } - if !found { - t.Errorf("/usage after turn: expected 'requests' in output, got: %v", byRole(evs, "info")) - } -} - -// --- /history --- - -func TestCommand_History_NoSession(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/history") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "no active session") { - found = true - } - } - if !found { - t.Errorf("/history before any turn: expected 'no active session', got: %v", byRole(evs, "info")) - } -} - -func TestCommand_History_WithSession(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "remember this") - evs := collectEvents(context.Background(), c, "/history") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "user") && strings.Contains(s, "remember this") { - found = true - } - } - if !found { - t.Errorf("/history: user message not found in output: %v", byRole(evs, "info")) - } -} - -// --- /sp --- - -func TestCommand_SP(t *testing.T) { - c := newCore(t, nil, nil) - evs := collectEvents(context.Background(), c, "/sp") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "test system prompt") { - found = true - } - } - if !found { - t.Errorf("/sp: system prompt not in info events: %v", byRole(evs, "info")) - } -} - -// --- /agents with files --- - -func TestCommand_Agents_List(t *testing.T) { - c := newCore(t, nil, nil) - t.Setenv("OLLIE_AGENTS_PATH", c.r.agentsDir) - if err := os.WriteFile(c.r.agentsDir+"/myagent.json", []byte(`{}`), 0600); err != nil { - t.Fatal(err) - } - evs := collectEvents(context.Background(), c, "/agents") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "myagent") { - found = true - } - } - if !found { - t.Errorf("/agents with files: 'myagent' not in output: %v", byRole(evs, "info")) - } -} - -// --- hooksRan plural --- - -func TestHooksRan(t *testing.T) { - if got := hooksRan(1); got != "1 hook run" { - t.Errorf("hooksRan(1) = %q; want %q", got, "1 hook run") - } - if got := hooksRan(3); got != "3 hooks run" { - t.Errorf("hooksRan(3) = %q; want %q", got, "3 hooks run") - } -} - -// --- HookResult.Summary --- - -func TestHookResult_Summary(t *testing.T) { - tests := []struct { - name string - hr HookResult - want string - }{ - {"not ran", HookResult{}, ""}, - {"all success", HookResult{Ran: true, Total: 2, Succeeded: 2}, "(2 of 2 hooks run)"}, - {"one failure", HookResult{Ran: true, Total: 3, Succeeded: 2, Failed: 1, FailedCmds: []string{"bad.sh"}}, "(3 of 3 hooks run) (1 of 3 failed: bad.sh)"}, - {"all failed", HookResult{Ran: true, Total: 2, Succeeded: 0, Failed: 2, FailedCmds: []string{"a.sh", "b.sh"}}, "(2 of 2 hooks run) (2 of 2 failed: a.sh, b.sh)"}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := tt.hr.Summary(); got != tt.want { - t.Errorf("Summary() = %q; want %q", got, tt.want) - } - }) - } -} - -func TestTruncateCmd(t *testing.T) { - short := "echo hello" - if got := truncateCmd(short); got != short { - t.Errorf("truncateCmd(%q) = %q; want %q", short, got, short) - } - long := "this is a very long command that exceeds forty characters total" - got := truncateCmd(long) - if len(got) != 40 || !strings.HasSuffix(got, "...") { - t.Errorf("truncateCmd(long) = %q; want 40 chars ending in ...", got) - } -} - -// --- rate-limit retry --- - -func TestRun_RateLimitRetry(t *testing.T) { - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - if callCount == 1 { - return nil, &backend.RateLimitError{RetryAfter: time.Millisecond} - } - return textStream("ok"), nil - } - c := newCore(t, be, nil) - var retryEvents []Event - sub := c.Bus().Subscribe("event", func(ev Event) { - if ev.Role == "retry" { - retryEvents = append(retryEvents, ev) - } - }) - c.Submit(context.Background(), "hello") - c.Bus().Unsubscribe(sub) - if callCount != 2 { - t.Errorf("backend called %d times; want 2 (retry + success)", callCount) - } - if len(retryEvents) == 0 { - t.Error("no retry events emitted during rate-limit retry") - } -} - -// --- run: tool cancelled before execution --- - -func TestRun_ToolCancelledBeforeExec(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - cancel() // cancel before tool loop runs - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{ - ToolCalls: []backend.ToolCall{{Name: "my_tool", Arguments: json.RawMessage(`{}`)}}, - Done: true, - StopReason: "tool_calls", - } - close(ch) - return ch, nil - - } - c := newCore(t, be, nil) - c.Submit(ctx, "hello") - if callCount != 1 { - t.Errorf("backend called %d times; want 1 (no retry after cancellation)", callCount) - } - if got := c.State(); got != "idle" { - t.Errorf("State() = %q after tool cancellation; want idle", got) - } -} - -// --- run: ChatStream error while ctx cancelled (recordInterruption "request") --- - -func TestRun_RequestCancelled(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - cancel() - return nil, fmt.Errorf("request failed") - } - c := newCore(t, be, nil) - c.Submit(ctx, "hello") - if got := c.State(); got != "idle" { - t.Errorf("State() = %q after request cancellation; want idle", got) - } -} - -// --- run: Exec error while ctx cancelled, with pending inject --- - -func TestRun_ExecCancelledWithInject(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{ - ToolCalls: []backend.ToolCall{{Name: "my_tool", Arguments: json.RawMessage(`{}`)}}, - Done: true, - StopReason: "tool_calls", - } - close(ch) - return ch, nil - - } - c := newCore(t, be, nil) - inject := "user interrupt" - c.r.pendingInject.Store(&inject) - c.r.runtime.Exec = func(_ context.Context, _ string, _ json.RawMessage) (string, []backend.ContentBlock, error) { - cancel() // ctx cancelled during exec - return "", nil, fmt.Errorf("exec cancelled") - } - c.Submit(ctx, "hello") - if callCount != 1 { - t.Errorf("backend called %d times; want 1", callCount) - } - if got := c.State(); got != "idle" { - t.Errorf("State() = %q after exec cancellation; want idle", got) - } -} - -// --- auto-compact with hook context injection --- - -func TestAutoCompact_WithHookContext(t *testing.T) { - callCount := 0 - be := &mockBackend{name: "mock", model: "test", ctxLen: 10} - c := newCore(t, be, Hooks{ - HookPreCompact: []string{`echo "pre-compact context"`}, - HookPostCompact: []string{`echo "post-compact context"`}, - }) - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - if callCount == 1 { - return textStream("summary text for compaction"), nil - - } - return textStream("answer"), nil - - } - c.r.history = newHistory("goal") - for range 15 { - c.r.history.messages = append(c.r.history.messages, - backend.Message{Role: "assistant", Content: "a long response that exceeds seven tokens of content"}, - backend.Message{Role: "user", Content: "follow up question with enough text to push over the limit"}, - ) - } - evs := collectEvents(context.Background(), c, "next prompt") - if callCount < 2 { - t.Errorf("backend called %d times; want ≥2 (compact + turn)", callCount) - } - // Verify auto-compact ran by checking for the info message - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "auto-compacting") { - found = true - } - } - if !found { - t.Errorf("no 'auto-compacting' info event; got: %v", byRole(evs, "info")) - } -} - -// --- Core method accessors --- - -func TestCore_AgentName(t *testing.T) { - c := newCore(t, nil, nil) - if got := c.AgentName(); got != "test" { - t.Errorf("AgentName() = %q; want %q", got, "test") - } -} - -func TestCore_BackendName(t *testing.T) { - c := newCore(t, nil, nil) - if got := c.BackendName(); got != "mock" { - t.Errorf("BackendName() = %q; want %q", got, "mock") - } -} - -func TestCore_ModelName(t *testing.T) { - c := newCore(t, nil, nil) - if got := c.ModelName(); got != "test" { - t.Errorf("ModelName() = %q; want %q", got, "test") - } -} - -func TestCore_SystemPrompt(t *testing.T) { - c := newCore(t, nil, nil) - if got := c.SystemPrompt(); got != "test system prompt" { - t.Errorf("SystemPrompt() = %q; want %q", got, "test system prompt") - } -} - -func TestCore_GenerationParams(t *testing.T) { - c := newCore(t, nil, nil) - p := c.GenerationParams() - if p.MaxTokens != 0 { - t.Errorf("GenerationParams().MaxTokens = %d; want 0 (default)", p.MaxTokens) - } -} - -func TestCore_SetGenerationParams_WhileRunning(t *testing.T) { - unblock := make(chan struct{}) - be := defaultBE() - be.respond = func(ctx context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return blockedStream(ctx, unblock), nil - - } - c := newCore(t, be, nil) - - done := make(chan struct{}) - go func() { - defer close(done) - c.Submit(context.Background(), "hello") - }() - waitState(t, c, "thinking") - - err := c.SetGenerationParams(backend.GenerationParams{}) - - close(unblock) - <-done - - if err == nil { - t.Error("SetGenerationParams while running: want error, got nil") - } -} - -func TestCore_ListModels(t *testing.T) { - be := &mockBackend{name: "mock", model: "test", ctxLen: 128000, models: []string{"a", "b"}} - c := newCore(t, be, nil) - got := c.ListModels() - if !strings.Contains(got, "a") || !strings.Contains(got, "b") { - t.Errorf("ListModels() = %q; want 'a' and 'b'", got) - } -} - -func TestCore_CtxSz_NoSession(t *testing.T) { - c := newCore(t, nil, nil) - if got := c.CtxSz(); got != "no active session" { - t.Errorf("CtxSz() with no session = %q; want 'no active session'", got) - } -} - -func TestCore_CtxSz_WithSession(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "hello") - got := c.CtxSz() - if !strings.Contains(got, "/") { - t.Errorf("CtxSz() = %q; want token fraction like '10 / 128000 (0%%)'", got) - } -} - -func TestCore_Usage_NoSession(t *testing.T) { - c := newCore(t, nil, nil) - if got := c.Usage(); got != "no active session" { - t.Errorf("Usage() with no session = %q; want 'no active session'", got) - } -} - -func TestCore_Usage_WithSession(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "hello") - got := c.Usage() - if !strings.Contains(got, "requests") { - t.Errorf("Usage() = %q; want string containing 'requests'", got) - } -} - -// --- SetSessionID --- - -func TestSetSessionID_Rename(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "hello") // saves session file - oldID := c.id - newID := NewSessionID() - if err := c.SetSessionID(newID); err != nil { - t.Fatalf("SetSessionID: %v", err) - } - if c.id != newID { - t.Errorf("sessionID = %q; want %q", c.id, newID) - } - if _, err := os.Stat(filepath.Join(c.sessionsDir, "active", oldID+".json")); !os.IsNotExist(err) { - t.Errorf("old active session file still exists after rename; err=%v", err) - } - if _, err := os.Stat(filepath.Join(c.sessionsDir, "active", newID+".json")); err != nil { - t.Errorf("new active session file not found after rename: %v", err) - } -} - -func TestSetSessionID_UpdatesPreamble(t *testing.T) { - oldID := NewSessionID() - env := &Runtime{ - Preamble: "session is " + oldID + " end", - } - c := New(Config{ - Backend: defaultBE(), - AgentName: "test", - AgentsDir: t.TempDir(), - SessionsDir: t.TempDir(), - SessionID: oldID, - CWD: t.TempDir(), - Runtime: env, - NewDispatcher: tools.NewDispatcher, - }) - t.Cleanup(c.Close) - - newID := NewSessionID() - if err := c.SetSessionID(newID); err != nil { - t.Fatalf("SetSessionID: %v", err) - } - want := "session is " + newID + " end" - if got := c.SystemPrompt(); got != want { - t.Errorf("SystemPrompt() = %q; want %q", got, want) - } -} - -func TestSetSessionID_SameID(t *testing.T) { - c := newCore(t, nil, nil) - id := c.id - if err := c.SetSessionID(id); err != nil { - t.Fatalf("SetSessionID with same ID: %v", err) - } - if c.id != id { - t.Errorf("sessionID changed: got %q; want %q", c.id, id) - } -} - -// --- SetGenerationParams success --- - -func TestCore_SetGenerationParams_Success(t *testing.T) { - c := newCore(t, nil, nil) - p := backend.GenerationParams{MaxTokens: 100} - if err := c.SetGenerationParams(p); err != nil { - t.Fatalf("SetGenerationParams: %v", err) - } - if got := c.GenerationParams().MaxTokens; got != 100 { - t.Errorf("GenerationParams().MaxTokens = %d; want 100", got) - } -} - -// --- autoCompactLimit with zero ctxLen --- - -func TestAutoCompactLimit_DefaultWhenZero(t *testing.T) { - be := &mockBackend{name: "mock", model: "test", ctxLen: 0} - c := newCore(t, be, nil) - limit := c.r.autoCompactLimit(context.Background()) - want := defaultContextLength * 3 / 4 - if limit != want { - t.Errorf("autoCompactLimit with ctxLen=0 = %d; want %d", limit, want) - } -} - -// --- agentSpawn hook with context output --- - -func TestAgentSpawn_WithContext(t *testing.T) { - c := newCore(t, nil, Hooks{HookAgentSpawn: []string{`echo "spawn context"`}}) - collectEvents(context.Background(), c, "first") - // spawn context goes into session history, not preamble - found := false - for _, m := range c.r.history.messages { - if strings.Contains(m.Content, "spawn context") { - found = true - break - } - } - if !found { - t.Errorf("spawn context not found in session history") - } -} - -// --- retryCountdown: ctx cancelled during wait --- - -func TestRun_RateLimitRetry_Cancelled(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - callCount := 0 - be := defaultBE() - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - go func() { - time.Sleep(10 * time.Millisecond) - cancel() - }() - return nil, &backend.RateLimitError{RetryAfter: 10 * time.Second} - } - c := newCore(t, be, nil) - c.Submit(ctx, "hello") - if callCount != 1 { - t.Errorf("backend called %d times after cancelled retry; want 1", callCount) - } -} - -// --- compact: backend error --- - -func TestManualCompact_BackendError(t *testing.T) { - callCount := 0 - be := defaultBE() - c := newCore(t, be, nil) - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ - return nil, fmt.Errorf("compact backend error") - } - c.r.history = newHistory("goal") - for i := range 15 { - c.r.history.messages = append(c.r.history.messages, - backend.Message{Role: "assistant", Content: fmt.Sprintf("response %d", i)}, - backend.Message{Role: "user", Content: fmt.Sprintf("follow up %d", i)}, - ) - } - evs := collectEvents(context.Background(), c, "/compact") - if callCount == 0 { - t.Fatal("backend not called for compact") - } - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "compact error") { - found = true - } - } - if !found { - t.Errorf("/compact backend error: expected 'compact error' event; got: %v", byRole(evs, "info")) - } -} - -// --- compact: empty summary from backend --- - -func TestManualCompact_EmptySummary(t *testing.T) { - be := defaultBE() - c := newCore(t, be, nil) - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - ch := make(chan backend.StreamEvent, 1) - ch <- backend.StreamEvent{Done: true, StopReason: "stop", Content: " "} - close(ch) - return ch, nil - - } - c.r.history = newHistory("goal") - for i := range 15 { - c.r.history.messages = append(c.r.history.messages, - backend.Message{Role: "assistant", Content: fmt.Sprintf("response %d", i)}, - backend.Message{Role: "user", Content: fmt.Sprintf("follow up %d", i)}, - ) - } - evs := collectEvents(context.Background(), c, "/compact") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "compact error") { - found = true - } - } - if !found { - t.Errorf("/compact empty summary: expected 'compact error'; got: %v", byRole(evs, "info")) - } -} - -// --- compact: session with tool-call messages (exercises flattenToolMessages) --- - -func TestManualCompact_WithToolMessages(t *testing.T) { - be := defaultBE() - c := newCore(t, be, nil) - be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - return textStream("summary"), nil - - } - c.r.history = newHistory("goal") - // Seed enough messages to pass the compaction threshold. - for i := range 10 { - c.r.history.messages = append(c.r.history.messages, - backend.Message{Role: "assistant", Content: fmt.Sprintf("response %d", i)}, - backend.Message{Role: "user", Content: fmt.Sprintf("follow up %d", i)}, - ) - } - c.r.history.messages = append(c.r.history.messages, - backend.Message{ - Role: "assistant", - Content: "calling tool", - ToolCalls: []backend.ToolCall{{Name: "my_tool", Arguments: json.RawMessage(`{"key":"val"}`)}}, - }, - backend.Message{Role: "tool", Content: "tool result", ToolCallID: "1"}, - backend.Message{Role: "user", Content: "more stuff"}, - backend.Message{Role: "assistant", Content: "done"}, - backend.Message{Role: "user", Content: "follow up"}, - ) - evs := collectEvents(context.Background(), c, "/compact") - found := false - for _, s := range byRole(evs, "info") { - if strings.Contains(s, "compacted") { - found = true - } - } - if !found { - t.Errorf("/compact with tool messages: no 'compacted' event; got: %v", byRole(evs, "info")) - } -} - -// --- hook timeout --- - -func TestHookTimeout_Branch(t *testing.T) { - old := hookTimeout - hookTimeout = 0 // 0s timeout fires immediately - t.Cleanup(func() { hookTimeout = old }) - - // A hook that sleeps longer than the timeout. The timeout branch kills the - // process and returns a warning, but does not block the turn. - hooks := Hooks{HookPreTurn: []string{"sleep 10"}} - c := newCore(t, nil, hooks) - evs := collectEvents(context.Background(), c, "hello") - - // Turn must have run: an assistant event proves the hook didn't block it. - if got := byRole(evs, "assistant"); len(got) == 0 { - t.Errorf("expected assistant event after hook timeout; hook must not have blocked the turn") - } - if c.State() != "idle" { - t.Errorf("State() = %q after timeout hook; want idle", c.State()) - } - // A timeout warning should have been emitted. - infos := byRole(evs, "info") - found := false - for _, s := range infos { - if strings.Contains(s, "timed out") { - found = true - } - } - if !found { - t.Errorf("expected 'timed out' warning in info events; got %v", infos) - } -} - -// --- /backend with injected newBackend --- - -func TestCommand_Backend_Switch(t *testing.T) { - c := newCore(t, nil, nil) - newBE := &mockBackend{name: "injected", model: "new-model"} - c.r.newBackend = func(string) (backend.Backend, error) { return newBE, nil } - - evs := collectEvents(context.Background(), c, "/backend other") - infos := byRole(evs, "info") - found := false - for _, s := range infos { - if strings.Contains(s, "injected") { - found = true - } - } - if !found { - t.Errorf("/backend switch: expected 'injected' in info events; got %v", infos) - } - if c.r.runtime.Backend != newBE { - t.Errorf("/backend switch: backend not updated") - } -} - -func TestCommand_Backend_Error(t *testing.T) { - c := newCore(t, nil, nil) - c.r.newBackend = func(string) (backend.Backend, error) { return nil, fmt.Errorf("no such backend") } - - evs := collectEvents(context.Background(), c, "/backend bad") - infos := byRole(evs, "info") - found := false - for _, s := range infos { - if strings.Contains(s, "no such backend") { - found = true - } - } - if !found { - t.Errorf("/backend error: expected error message; got %v", infos) - } -} - -// --- extractToolResult --- - -func TestExtractToolResult_Success(t *testing.T) { - raw := json.RawMessage(`{"isError":false,"content":[{"type":"text","text":"hello"}]}`) - text, _, isErr := extractToolResult(raw) - if text != "hello" { - t.Errorf("text = %q; want %q", text, "hello") - } - if isErr { - t.Error("isError should be false") - } -} - -func TestExtractToolResult_IsError(t *testing.T) { - raw := json.RawMessage(`{"isError":true,"content":[{"type":"text","text":"something failed"}]}`) - text, _, isErr := extractToolResult(raw) - if text != "something failed" { - t.Errorf("text = %q; want %q", text, "something failed") - } - if !isErr { - t.Error("isError should be true") - } -} - -func TestExtractToolResult_MultipleContentItems(t *testing.T) { - raw := json.RawMessage(`{"isError":false,"content":[{"type":"text","text":"a"},{"type":"text","text":"b"}]}`) - text, _, _ := extractToolResult(raw) - if text != "a\nb" { - t.Errorf("text = %q; want %q", text, "a\nb") - } -} - -func TestExtractToolResult_NonTextItemsSkipped(t *testing.T) { - raw := json.RawMessage(`{"isError":false,"content":[{"type":"image","text":"ignored"},{"type":"text","text":"kept"}]}`) - text, _, _ := extractToolResult(raw) - if text != "kept" { - t.Errorf("text = %q; want %q", text, "kept") - } -} - -func TestExtractToolResult_ImageBlock(t *testing.T) { - raw := json.RawMessage(`{"isError":false,"content":[{"type":"image","media_type":"image/png","data":"iVBOR"},{"type":"text","text":"screenshot above"}]}`) - text, blocks, isErr := extractToolResult(raw) - if text != "screenshot above" { - t.Errorf("text = %q; want %q", text, "screenshot above") - } - if isErr { - t.Error("isError should be false") - } - if len(blocks) != 1 { - t.Fatalf("blocks len = %d; want 1", len(blocks)) - } - b := blocks[0] - if b.Type != "image" || b.ImageSource == nil || b.ImageSource.MediaType != "image/png" || b.ImageSource.Data != "iVBOR" { - t.Errorf("block = %+v", b) - } -} - -func TestExtractToolResult_InvalidJSON(t *testing.T) { - raw := json.RawMessage(`not json`) - text, _, isErr := extractToolResult(raw) - if text != "not json" { - t.Errorf("text = %q; want raw input on parse failure", text) - } - if isErr { - t.Error("isError should be false on parse failure") - } -} - -// --- toolInfosToBackend --- - -func TestToolInfosToBackend(t *testing.T) { - schema := json.RawMessage(`{"type":"object"}`) - infos := []tools.ToolInfo{ - {Name: "tool_a", Description: "does A.", InputSchema: schema}, - {Name: "tool_b", Description: "does B.", InputSchema: schema}, - } - got := toolInfosToBackend(infos) - if len(got) != 2 { - t.Fatalf("len = %d; want 2", len(got)) - } - if got[0].Name != "tool_a" || got[0].Description != "does A." { - t.Errorf("got[0] = %+v", got[0]) - } - if string(got[1].Parameters) != string(schema) { - t.Errorf("got[1].Parameters = %s; want %s", got[1].Parameters, schema) - } -} - -func TestToolInfosToBackend_Empty(t *testing.T) { - got := toolInfosToBackend(nil) - if len(got) != 0 { - t.Errorf("expected empty slice, got %v", got) - } -} - -// --- RestoreSession round-trip --- - -func TestRestoreSession_RoundTrip(t *testing.T) { - dir := t.TempDir() - path := filepath.Join(dir, "sess.json") - - s := newHistory("first user message") - s.appendUserMessage("first user message") - s.appendUserMessage("second message") - s.messages = append(s.messages, backend.Message{Role: "assistant", Content: "reply"}) - - if err := s.saveTo(path, "test-id", "test-agent", "/tmp/test"); err != nil { - t.Fatalf("saveTo: %v", err) - } - - data, err := os.ReadFile(path) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - var ps PersistedAgent - if err := json.Unmarshal(data, &ps); err != nil { - t.Fatalf("Unmarshal: %v", err) - } - if ps.ID != "test-id" || ps.Agent != "test-agent" { - t.Errorf("ps.ID=%q ps.Agent=%q", ps.ID, ps.Agent) - } - - restored := RestoreHistory(&ps) - if restored.goal != "first user message" { - t.Errorf("goal = %q; want %q", restored.goal, "first user message") - } - if len(restored.messages) != len(s.messages) { - t.Errorf("messages len = %d; want %d", len(restored.messages), len(s.messages)) - } -} - -func TestReactionTargetsResponseAndReplaces(t *testing.T) { - c := newCore(t, nil, nil) - c.r.history = &History{messages: []backend.Message{ - {Role: "assistant", ID: "r1", Content: "first"}, - {Role: "assistant", ID: "r2", Content: "second"}, - }} - if err := c.ReactTo("r1", "👍"); err != nil { - t.Fatalf("ReactTo: %v", err) - } - if len(c.r.history.Reactions) != 1 || c.r.history.Reactions[0].ResponseID != "r1" || c.r.history.PositiveReactions != 1 { - t.Fatalf("reaction state = %+v positive=%d", c.r.history.Reactions, c.r.history.PositiveReactions) - } - if err := c.ReactTo("r1", "👎"); err != nil { - t.Fatalf("replace ReactTo: %v", err) - } - if len(c.r.history.Reactions) != 1 || c.r.history.PositiveReactions != 0 || c.r.history.NegativeReactions != 1 { - t.Fatalf("replaced reaction state = %+v positive=%d negative=%d", c.r.history.Reactions, c.r.history.PositiveReactions, c.r.history.NegativeReactions) - } - ctx := c.Context() - if got := ctx[len(ctx)-1].Content; !strings.Contains(got, "response r1: negative") { - t.Fatalf("reaction context = %q", got) - } - if err := c.ReactTo("missing", "👍"); err == nil { - t.Fatal("ReactTo missing response succeeded") - } -} - -func TestReactionPersistsAndRestores(t *testing.T) { - path := filepath.Join(t.TempDir(), "session.json") - s := &History{ - messages: []backend.Message{{Role: "assistant", ID: "r1", Content: "answer"}}, - Reactions: []Reaction{{ID: "x1", ResponseID: "r1", Emoji: "🚀", Category: "excellent", CreatedAt: time.Now()}}, - PositiveReactions: 1, - } - if err := s.saveTo(path, "id", "agent", "/tmp"); err != nil { - t.Fatal(err) - } - ps, err := LoadPersistedAgent(path) - if err != nil { - t.Fatal(err) - } - restored := RestoreHistory(ps) - if len(restored.Reactions) != 1 || restored.Reactions[0].ResponseID != "r1" || restored.PositiveReactions != 1 { - t.Fatalf("restored = %+v positive=%d", restored.Reactions, restored.PositiveReactions) - } -} - -func TestRestoreSession_GoalFromFirstUserMessage(t *testing.T) { - msgs := []backend.Message{ - {Role: "assistant", Content: "preamble"}, - {Role: "user", Content: "the real goal"}, - {Role: "user", Content: "second user msg"}, - } - s := RestoreHistory(&PersistedAgent{Messages: msgs}) - if s.goal != "the real goal" { - t.Errorf("goal = %q; want %q", s.goal, "the real goal") - } -} - -// --- SetEnv propagation --- - -// mockEnvServer is a tools.Server that also implements tools.EnvSetter. -type mockEnvServer struct { - mu sync.Mutex - env map[string]string -} - -func (m *mockEnvServer) ListTools() ([]tools.ToolInfo, error) { return nil, nil } -func (m *mockEnvServer) CallTool(_ context.Context, _ string, _ json.RawMessage) (json.RawMessage, error) { - return nil, fmt.Errorf("not implemented") -} -func (m *mockEnvServer) SetEnv(k, v string) { - m.mu.Lock() - defer m.mu.Unlock() - if m.env == nil { - m.env = make(map[string]string) - } - m.env[k] = v -} -func (m *mockEnvServer) get(k string) string { - m.mu.Lock() - defer m.mu.Unlock() - return m.env[k] -} - -func newCoreWithExecServer(t *testing.T, srv *mockEnvServer) *Session { - t.Helper() - t.Setenv("OLLIE", "") - d := tools.NewDispatcher() - d.AddServer("execute", srv) - env := &Runtime{ - Hooks: Hooks{}, - Preamble: "test system prompt", - Dispatcher: d, - } - c := New(Config{ - Backend: defaultBE(), - AgentName: "test", - AgentsDir: t.TempDir(), - SessionsDir: t.TempDir(), - SessionID: NewSessionID(), - CWD: t.TempDir(), - Runtime: env, - NewDispatcher: tools.NewDispatcher, - }) - t.Cleanup(c.Close) - return c -} - -func TestSetEnv_PropagatestoExecuteServer(t *testing.T) { - srv := &mockEnvServer{} - c := newCoreWithExecServer(t, srv) - - c.SetEnv("MY_KEY", "my_value") - - if got := srv.get("MY_KEY"); got != "my_value" { - t.Errorf("execute server env MY_KEY = %q; want %q", got, "my_value") - } -} - -func TestSetEnv_StoredInCore(t *testing.T) { - srv := &mockEnvServer{} - c := newCoreWithExecServer(t, srv) - - c.SetEnv("FOO", "bar") - c.SetEnv("BAZ", "qux") - - c.envMu.RLock() - env := c.env - c.envMu.RUnlock() - if env["FOO"] != "bar" { - t.Errorf("env[FOO] = %q; want bar", env["FOO"]) - } - if env["BAZ"] != "qux" { - t.Errorf("env[BAZ] = %q; want qux", env["BAZ"]) - } -} - -func TestSetEnv_NilDispatcher_NoPanic(t *testing.T) { - c := newCore(t, nil, nil) - c.r.runtime.Dispatcher = nil - // Must not panic. - c.SetEnv("K", "V") - c.envMu.RLock() - env := c.env - c.envMu.RUnlock() - if env["K"] != "V" { - t.Errorf("env[K] = %q; want V", env["K"]) - } -} - -// --- firstSentence --- - -func TestFirstSentence_Period(t *testing.T) { - if got := firstSentence("Does a thing. More detail."); got != "Does a thing." { - t.Errorf("got %q", got) - } -} - -func TestFirstSentence_Newline(t *testing.T) { - if got := firstSentence("Does a thing\nMore detail"); got != "Does a thing" { - t.Errorf("got %q", got) - } -} - -func TestFirstSentence_TruncatesLong(t *testing.T) { - long := strings.Repeat("x", 100) - got := firstSentence(long) - if len(got) != 80 || !strings.HasSuffix(got, "...") { - t.Errorf("got %q (len %d)", got, len(got)) - } -} - -func TestFirstSentence_ShortNoSentenceEnd(t *testing.T) { - if got := firstSentence("short"); got != "short" { - t.Errorf("got %q", got) - } -} - -// --- BuildRuntime (nil config path) --- - -func setupCfgDir(t *testing.T) string { - t.Helper() - dir := t.TempDir() - t.Setenv("OLLIE_CFG_PATH", dir) - t.Setenv("OLLIE_TOOLS_PATH", filepath.Join(dir, "tools")) - return dir -} - -func TestBuildRuntime_NilConfig(t *testing.T) { - setupCfgDir(t) - d := tools.NewDispatcher() - env := BuildRuntime(nil, d, t.TempDir(), nil) - - if env.Preamble != "" { - t.Errorf("preamble = %q; want empty", env.Preamble) - } - if len(env.Hooks) != 0 { - t.Errorf("expected no hooks; got %v", env.Hooks) - } - if len(env.Tools) != 0 { - t.Errorf("expected no tools; got %v", env.Tools) - } - if len(env.Messages) != 0 { - t.Errorf("expected no startup messages; got %v", env.Messages) - } -} - -func TestBuildRuntime_PromptBecomesPreamble(t *testing.T) { - setupCfgDir(t) - d := tools.NewDispatcher() - cfg := &AgentConfig{Prompt: Prompt{Value: []string{"the prompt"}}} - env := BuildRuntime(cfg, d, t.TempDir(), nil) - - if env.Preamble != "the prompt" { - t.Errorf("preamble = %q; want %q", env.Preamble, "the prompt") - } -} - -func TestBuildRuntime_HooksAndParams(t *testing.T) { - setupCfgDir(t) - d := tools.NewDispatcher() - temp := 0.7 - cfg := &AgentConfig{ - Hooks: map[string]HookCmds{"preTurn": {"echo hi"}}, - MaxTokens: 512, - Temperature: &temp, - } - env := BuildRuntime(cfg, d, t.TempDir(), nil) - - if cmds := env.Hooks[HookPreTurn]; len(cmds) != 1 || cmds[0] != "echo hi" { - t.Errorf("Hooks[preTurn] = %v", cmds) - } - if env.GenParams.MaxTokens != 512 { - t.Errorf("MaxTokens = %d; want 512", env.GenParams.MaxTokens) - } - if env.GenParams.Temperature == nil || *env.GenParams.Temperature != 0.7 { - t.Errorf("Temperature = %v; want 0.7", env.GenParams.Temperature) - } -} - -func TestBuildRuntime_ExecUnknownTool(t *testing.T) { - setupCfgDir(t) - d := tools.NewDispatcher() - env := BuildRuntime(nil, d, t.TempDir(), nil) - - _, _, err := env.Exec(context.Background(), "no_such_tool", json.RawMessage(`{}`)) - if err == nil || !strings.Contains(err.Error(), "unknown tool") { - t.Errorf("expected unknown tool error; got %v", err) - } -} - -// --- mock dispatcher for BuildRuntime tests --- - -type mockDispatcher struct { - tools []tools.ToolInfo - listErr error - servers map[string]tools.Server - dispatch func(ctx context.Context, server, tool string, args json.RawMessage) (json.RawMessage, error) -} - -func (m *mockDispatcher) AddServer(name string, s tools.Server) { - if m.servers == nil { - m.servers = make(map[string]tools.Server) - } - m.servers[name] = s -} -func (m *mockDispatcher) GetServer(name string) (tools.Server, bool) { - s, ok := m.servers[name] - return s, ok -} -func (m *mockDispatcher) ListTools() ([]tools.ToolInfo, error) { - return m.tools, m.listErr -} -func (m *mockDispatcher) Dispatch(ctx context.Context, server, tool string, args json.RawMessage) (json.RawMessage, error) { - if m.dispatch != nil { - return m.dispatch(ctx, server, tool, args) - } - return nil, fmt.Errorf("not implemented") -} - -func TestBuildRuntime_ListToolsError(t *testing.T) { - setupCfgDir(t) - d := &mockDispatcher{listErr: fmt.Errorf("boom")} - env := BuildRuntime(nil, d, t.TempDir(), nil) - if len(env.Messages) != 1 || !strings.Contains(env.Messages[0], "boom") { - t.Errorf("Messages = %v", env.Messages) - } -} - -func TestBuildRuntime_ToolsPopulated(t *testing.T) { - setupCfgDir(t) - d := &mockDispatcher{ - tools: []tools.ToolInfo{{Server: "s1", Name: "mytool", Description: "desc", InputSchema: json.RawMessage(`{}`)}}, - } - env := BuildRuntime(nil, d, t.TempDir(), nil) - if len(env.Tools) != 1 || env.Tools[0].Name != "mytool" { - t.Errorf("tools = %+v", env.Tools) - } -} - -func TestBuildRuntime_ToolsDisabled(t *testing.T) { - setupCfgDir(t) - d := &mockDispatcher{ - tools: []tools.ToolInfo{{Server: "s1", Name: "mytool", Description: "desc", InputSchema: json.RawMessage(`{}`)}}, - } - f := false - cfg := &AgentConfig{Tools: &f} - env := BuildRuntime(cfg, d, t.TempDir(), nil) - if len(env.Tools) != 0 { - t.Errorf("expected no tools when disabled; got %+v", env.Tools) - } -} - -func TestBuildRuntime_PromptOnly(t *testing.T) { - setupCfgDir(t) - d := tools.NewDispatcher() - cfg := &AgentConfig{Prompt: Prompt{Value: []string{"only agent"}}} - env := BuildRuntime(cfg, d, t.TempDir(), nil) - if env.Preamble != "only agent" { - t.Errorf("preamble = %q; want %q", env.Preamble, "only agent") - } -} - -func TestBuildRuntime_ExecPrompt(t *testing.T) { - setupCfgDir(t) - d := tools.NewDispatcher() - cfg := &AgentConfig{Prompt: Prompt{ - Value: []string{"echo hello", "echo 'You are a security auditor.'", "echo world"}, - IsExec: true, - }} - env := BuildRuntime(cfg, d, t.TempDir(), nil) - want := "hello\nYou are a security auditor.\nworld" - if env.Preamble != want { - t.Errorf("preamble = %q; want %q", env.Preamble, want) - } -} - -func TestBuildRuntime_ExecPromptFileResolution(t *testing.T) { - dir := setupCfgDir(t) - // Create a prompts directory with a test prompt file. - promptsDir := filepath.Join(dir, "prompts") - os.MkdirAll(promptsDir, 0o755) - os.WriteFile(filepath.Join(promptsDir, "test-prompt.md"), []byte("Hello ${PRIME_PLATFORM}"), 0o644) - t.Setenv("OLLIE_PROMPTS_PATH", promptsDir) - - d := tools.NewDispatcher() - cfg := &AgentConfig{Prompt: Prompt{ - Value: []string{"test-prompt", "echo extra"}, - IsExec: true, - }} - cwd := t.TempDir() - rt := BuildRuntime(cfg, d, cwd, PromptEnv(cwd)) - if !strings.Contains(rt.Preamble, "Hello linux") { - t.Errorf("expected file content with expanded vars, got %q", rt.Preamble) - } - if !strings.Contains(rt.Preamble, "extra") { - t.Errorf("expected shell fallback output, got %q", rt.Preamble) - } -} - -func TestBuildRuntime_ExecDispatchSuccess(t *testing.T) { - setupCfgDir(t) - d := &mockDispatcher{ - tools: []tools.ToolInfo{{Server: "s1", Name: "mytool"}}, - dispatch: func(ctx context.Context, server, tool string, args json.RawMessage) (json.RawMessage, error) { - return json.RawMessage(`{"content":[{"type":"text","text":"ok"}]}`), nil - }, - } - env := BuildRuntime(nil, d, t.TempDir(), nil) - result, _, err := env.Exec(context.Background(), "mytool", json.RawMessage(`{}`)) - if err != nil { - t.Fatal(err) - } - if result != "ok" { - t.Errorf("result = %q", result) - } -} - -func TestBuildRuntime_ExecDispatchError(t *testing.T) { - setupCfgDir(t) - d := &mockDispatcher{ - tools: []tools.ToolInfo{{Server: "s1", Name: "mytool"}}, - dispatch: func(ctx context.Context, server, tool string, args json.RawMessage) (json.RawMessage, error) { - return nil, fmt.Errorf("dispatch failed") - }, - } - env := BuildRuntime(nil, d, t.TempDir(), nil) - _, _, err := env.Exec(context.Background(), "mytool", json.RawMessage(`{}`)) - if err == nil || !strings.Contains(err.Error(), "dispatch failed") { - t.Errorf("err = %v", err) - } -} - -func TestBuildRuntime_ExecToolResultIsError(t *testing.T) { - setupCfgDir(t) - d := &mockDispatcher{ - tools: []tools.ToolInfo{{Server: "s1", Name: "mytool"}}, - dispatch: func(ctx context.Context, server, tool string, args json.RawMessage) (json.RawMessage, error) { - return json.RawMessage(`{"isError":true,"content":[{"type":"text","text":"bad thing"}]}`), nil - }, - } - env := BuildRuntime(nil, d, t.TempDir(), nil) - _, _, err := env.Exec(context.Background(), "mytool", json.RawMessage(`{}`)) - if err == nil || !strings.Contains(err.Error(), "bad thing") { - t.Errorf("err = %v", err) - } -} diff --git a/session/cost_test.go b/session/cost_test.go deleted file mode 100644 index 3e7f6e6..0000000 --- a/session/cost_test.go +++ /dev/null @@ -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) - } -} diff --git a/session/loop_cache_test.go b/session/loop_cache_test.go deleted file mode 100644 index 915c7cf..0000000 --- a/session/loop_cache_test.go +++ /dev/null @@ -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 } diff --git a/session/loop_errlimit_test.go b/session/loop_errlimit_test.go deleted file mode 100644 index bfcbe4e..0000000 --- a/session/loop_errlimit_test.go +++ /dev/null @@ -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") - } -} diff --git a/session/loop_maxsteps_helpers_test.go b/session/loop_maxsteps_helpers_test.go deleted file mode 100644 index aa5aea1..0000000 --- a/session/loop_maxsteps_helpers_test.go +++ /dev/null @@ -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) {} diff --git a/session/loop_maxsteps_test.go b/session/loop_maxsteps_test.go deleted file mode 100644 index b471048..0000000 --- a/session/loop_maxsteps_test.go +++ /dev/null @@ -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 } diff --git a/session/loop_parallel_test.go b/session/loop_parallel_test.go deleted file mode 100644 index a62a4c6..0000000 --- a/session/loop_parallel_test.go +++ /dev/null @@ -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 - } - } -} diff --git a/session/loop_truncate_test.go b/session/loop_truncate_test.go deleted file mode 100644 index 74403fe..0000000 --- a/session/loop_truncate_test.go +++ /dev/null @@ -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)]) - } - } -} diff --git a/session/loop_turnerror_test.go b/session/loop_turnerror_test.go deleted file mode 100644 index 0cd4e39..0000000 --- a/session/loop_turnerror_test.go +++ /dev/null @@ -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) - } - } -} diff --git a/session/session.go b/session/session.go index 85bf4c5..8422e3d 100644 --- a/session/session.go +++ b/session/session.go @@ -3,13 +3,10 @@ package session import ( "context" "crypto/rand" - "encoding/json" - "errors" "fmt" "io" "os" "path/filepath" - "slices" "strconv" "strings" "sync" @@ -17,385 +14,66 @@ import ( "time" "github.com/simonfxr/pubsub" + "ollie/agent" "ollie/backend" olog "ollie/log" "ollie/paths" "ollie/tools" ) -// 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 - -// 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 - serverOf := make(map[string]string) - - if cfg == nil || cfg.ToolsEnabled() { - var listErr error - allToolInfos, listErr = d.ListTools() - if listErr != nil { - messages = append(messages, fmt.Sprintf("list tools: %v", listErr)) - } - for _, t := range allToolInfos { - serverOf[t.Name] = t.Server - } - // Only built-in executors (with InputSchema) become backend tools. - // Named tool scripts are promoted via the tool registry and appear in the preamble. - 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) { - // Dynamically resolve the server for this tool — handles - // lazy-promoted tools that appeared after BuildRuntime. - 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). - // Full prompts are available on-demand via /tools write (net/dns pattern). - var toolListing strings.Builder - for _, ti := range allToolInfos { - if ti.Description != "" && ti.Server == "" { - // Only named tool scripts (Server==""), not built-in executors - 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, - } +// Config is the configuration for creating a session. +type Config struct { + Backend backend.Backend + ModelName string + AgentName string + AgentsDir string + SessionsDir string + SessionID string + AgentID string + CWD string + History *agent.History + Runtime *agent.Runtime + NewDispatcher func() tools.Dispatcher + NewBackend func(string) (backend.Backend, error) + Log *olog.Logger + MaxSteps int + ReadPlanStep func() string + ListHandlers map[string]func() []string + PromptEnvExtra []string + Remote string + BaseLayers []string } -// DefaultPromptsDir returns the default directory for prompt templates. -func DefaultPromptsDir() string { - return paths.CfgDir() + "/prompts" -} +// Session is the concrete session type. It owns session-level state and +// delegates agent operations to its owned Agent. +type Session struct { + id string + bus *pubsub.Bus + envMu sync.RWMutex + env map[string]string + plan []byte + prevPrompt string -// 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()} -} + r *agent.Agent + log *olog.Logger + sessionsDir string + listHandlers map[string]func() []string + remote string + mu sync.RWMutex + auditLog *olog.Logger -// 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" + saveMu sync.Mutex + saveDirty bool + saveTimer *time.Timer } // NewSessionID generates a unique, lexicographically sortable session identifier. -// Format: - — sortable by creation time, unique -// even if two sessions are created within the same nanosecond. func NewSessionID() string { b := make([]byte, 3) rand.Read(b) //nolint:errcheck return strconv.FormatInt(time.Now().UnixNano(), 10) + "-" + fmt.Sprintf("%06x", b) } -// 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) -} - -// infoEvent wraps a plain-text message as an info Event. -func infoEvent(text string) Event { - return Event{Role: "info", Content: text + "\n"} -} - -// actionHandle holds the cancel function for the current agent turn. -type actionHandle struct { - cancel context.CancelCauseFunc -} - -// Config is the configuration for creating an agent. -type Config struct { - Backend backend.Backend - ModelName string // if non-empty, overrides backend's default model - AgentName string - AgentsDir string - SessionsDir string - SessionID string - AgentID string // unique agent identity - CWD string // working directory for tool execution and system prompt - History *History - Runtime *Runtime - NewDispatcher func() tools.Dispatcher - NewBackend func(string) (backend.Backend, error) // if nil, defaults to backend.NewWithName - Log *olog.Logger // if nil, logging is disabled - // MaxSteps overrides the agent JSON maxSteps when non-zero. - // 0 means use the value from the agent config (or unlimited if absent). - MaxSteps int - // ReadPlanStep, if non-nil, is called to get the next unchecked plan step. - // When provided, it is used directly instead of reading from the filesystem. - ReadPlanStep func() string - // ListHandlers provides session-specific list implementations for slash - // commands such as /skills, /tools, /agents. Keyed by command name without - // the leading slash. Commands with no entry emit nothing. - ListHandlers map[string]func() []string - // PromptEnvExtra holds PRIME_* env vars for prompt resolution. - // For remote sessions these come from HostInfo; for local, from PromptEnv(). - // Stored so /agent reloads use the correct values. - PromptEnvExtra []string - // Remote is the SSH target for remote execution (empty = local). - Remote string - // BaseLayers are the system prompt, operational model, and environment - // block prepended to the agent preamble. Stored so /agent reloads - // preserve the server-injected context. - BaseLayers []string -} - -// Session is the concrete session type. It owns all session and agent state -// but has no knowledge of how output is rendered. -type Session struct { - // Session-level state - id string - bus *pubsub.Bus - envMu sync.RWMutex - env map[string]string - plan []byte - prevPrompt string - - // Agent - r *Agent - log *olog.Logger - sessionsDir string - listHandlers map[string]func() []string - remote string // SSH target for remote execution - mu sync.RWMutex - auditLog *olog.Logger - - // Debounced session persistence - saveMu sync.Mutex - saveDirty bool - saveTimer *time.Timer - -} - -// ToolCallCount returns the total number of tool calls executed in this -// session. The counter is monotonically increasing and never resets. -// Blocked calls (pre-tool hook exit 2) are not counted. -func (a *Session) ToolCallCount() int64 { - return a.r.toolCallCount.Load() -} - -// SetEnv stores a session-scoped variable and propagates it to the execute server. -func (a *Session) SetEnv(key, value string) { - a.envMu.Lock() - a.env[key] = value - a.envMu.Unlock() - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil { - return - } - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if es, ok := srv.(tools.EnvSetter); ok { - es.SetEnv(key, value) - } - } -} - -// pushSessionEnv injects OLLIE_SESSION_ID into the execute server subprocess env. -func (a *Session) pushSessionEnv() { - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil || a.id == "" { - return - } - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if es, ok := srv.(tools.EnvSetter); ok { - es.SetEnv("OLLIE_SESSION_ID", a.id) - if a.r.id != "" { - es.SetEnv("OLLIE_UNAME", a.r.id) - } - } - } -} - -// pushLockDir sets the flock directory on the execute server to the session tmpdir. -func (a *Session) pushLockDir() { - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil || a.id == "" { - return - } - -} - - -var sweepTmpOnce sync.Once - -// ollieTmpDir returns the base temp directory for session tmpdirs. -func ollieTmpDir() string { - if p := os.Getenv("OLLIE_TMP_PATH"); p != "" { - return p - } - return filepath.Join(os.TempDir(), "ollie") -} - -// readNextUnchecked reads a plan file and returns the first unchecked step. -func readNextUnchecked(path string) string { - data, err := os.ReadFile(path) - if err != nil || len(data) == 0 { - return "" - } - return NextUncheckedStep(data) -} - // NextUncheckedStep returns the first unchecked step from plan bytes. func NextUncheckedStep(data []byte) string { for _, line := range strings.Split(string(data), "\n") { @@ -407,9 +85,15 @@ func NextUncheckedStep(data []byte) string { return "" } -// sweepStaleTmpDirs removes session tmpdirs left by a previous crash. -// All session tmpdirs are owned by a single process, so anything -// present at startup is stale. +var sweepTmpOnce sync.Once + +func ollieTmpDir() string { + if p := os.Getenv("OLLIE_TMP_PATH"); p != "" { + return p + } + return filepath.Join(os.TempDir(), "ollie") +} + func sweepStaleTmpDirs() { sweepTmpOnce.Do(func() { base := ollieTmpDir() @@ -418,7 +102,7 @@ func sweepStaleTmpDirs() { }) } -// New creates an agent from the given configuration. +// New creates a session with an owned agent from the given configuration. func New(cfg Config) *Session { sweepStaleTmpDirs() if cfg.ModelName != "" { @@ -429,22 +113,15 @@ func New(cfg Config) *Session { } rt := cfg.Runtime if rt == nil { - rt = &Runtime{} + rt = &agent.Runtime{} } - // Store the backend on the runtime so it's the single source of truth. rt.Backend = cfg.Backend - // A non-zero MaxSteps in Config takes precedence over the - // value loaded from the agent JSON. if cfg.MaxSteps > 0 { rt.MaxSteps = cfg.MaxSteps } - var readPlanStep func() string if cfg.SessionID != "" { os.MkdirAll(filepath.Join(ollieTmpDir(), cfg.SessionID), 0700) //nolint:errcheck } - if cfg.ReadPlanStep != nil { - readPlanStep = cfg.ReadPlanStep - } log := cfg.Log if log == nil { @@ -455,362 +132,154 @@ func New(cfg Config) *Session { auditLog := log.Sub("audit") a := &Session{ - id: cfg.SessionID, - bus: bus, - env: make(map[string]string), - r: &Agent{ - history: cfg.History, - runtime: rt, - cwd: paths.ExpandHome(cfg.CWD), - id: cfg.AgentID, - agentName: cfg.AgentName, - agentsDir: cfg.AgentsDir, - promptEnvExtra: cfg.PromptEnvExtra, - baseLayers: cfg.BaseLayers, - newDispatcher: cfg.NewDispatcher, - newBackend: cfg.NewBackend, - bus: bus, - log: log, - auditLog: auditLog, - sessionID: cfg.SessionID, - startupMessages: rt.Messages, - readPlanStep: readPlanStep, - }, + id: cfg.SessionID, + bus: bus, + env: make(map[string]string), log: log, auditLog: auditLog, sessionsDir: cfg.SessionsDir, remote: cfg.Remote, listHandlers: cfg.ListHandlers, } - a.r.InitCond() - a.r.state = "idle" - a.r.flushSave = a.flushSave - a.r.saveSession = a.saveSession - a.r.turnError = a.r.defaultTurnError - a.pushSessionEnv() - a.pushLockDir() + + a.r = agent.NewAgent(agent.AgentCfg{ + History: cfg.History, + Runtime: rt, + AgentName: cfg.AgentName, + AgentsDir: cfg.AgentsDir, + AgentID: cfg.AgentID, + CWD: paths.ExpandHome(cfg.CWD), + BaseLayers: cfg.BaseLayers, + PromptEnvExtra: cfg.PromptEnvExtra, + NewDispatcher: cfg.NewDispatcher, + NewBackend: cfg.NewBackend, + Bus: bus, + Log: log, + AuditLog: auditLog, + SessionID: cfg.SessionID, + StartupMsgs: rt.Messages, + ReadPlanStep: cfg.ReadPlanStep, + SaveSession: a.saveSession, + FlushSave: a.flushSave, + }) + + a.r.SetSessionEnv(a.id) return a } -// Close releases resources for this session, including its tmpdir. +// Close releases resources for this session. func (a *Session) Close() { a.log.Debug("Close() session=%q", a.id) a.flushSave() - if a.r.runtime != nil && a.r.runtime.Dispatcher != nil { - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if c, ok := srv.(interface{ Close() }); ok { - a.log.Debug("Close() calling execute.Close()") - c.Close() - } - } - } + a.r.Close() if a.id != "" { os.RemoveAll(filepath.Join(ollieTmpDir(), a.id)) //nolint:errcheck } } -// execServer returns the execute server if available, or nil. -func (a *Session) execServer() interface{} { - if a.r.runtime == nil || a.r.runtime.Dispatcher == nil { - return nil - } - srv, _ := a.r.runtime.Dispatcher.GetServer("execute") - return srv +// SetEnv stores a session-scoped variable and propagates it to the agent. +func (a *Session) SetEnv(key, value string) { + a.envMu.Lock() + a.env[key] = value + a.envMu.Unlock() + a.r.SetEnv(key, value) } -func (a *Session) Detach() bool { - if srv := a.execServer(); srv != nil { - if d, ok := srv.(interface{ Detach() bool }); ok { - return d.Detach() - } +// Submit processes user input: dispatches commands, delegates turns to agent. +func (a *Session) Submit(ctx context.Context, input string) { + a.log.Debug("Submit() input_len=%d running=%v", len(input), a.IsRunning()) + if input == "" { + return } - return false + if a.handleCommand(ctx, input) { + return + } + a.r.Submit(ctx, input) } -func (a *Session) ListDetached() []DetachedInfo { - if srv := a.execServer(); srv != nil { - type listDetacher interface { - ListDetachedRaw() []any - } - if ld, ok := srv.(listDetacher); ok { - raw := ld.ListDetachedRaw() - out := make([]DetachedInfo, 0, len(raw)) - for _, r := range raw { - if m, ok := r.(map[string]any); ok { - di := DetachedInfo{} - if v, ok := m["pid"].(int); ok { - di.PID = v - } - if v, ok := m["command"].(string); ok { - di.Command = v - } - if v, ok := m["started"].(int64); ok { - di.Started = v - } - if v, ok := m["exited"].(bool); ok { - di.Exited = v - } - if v, ok := m["exit_code"].(int); ok { - di.ExitCode = v - } - out = append(out, di) - } - } - return out - } +func (a *Session) IsRunning() bool { return a.r.IsRunning() } +func (a *Session) State() string { return a.r.State() } +func (a *Session) Reply() string { return a.r.Reply() } +func (a *Session) AgentName() string { return a.r.Name() } +func (a *Session) BackendName() string { return a.r.BackendName() } +func (a *Session) ModelName() string { return a.r.ModelName() } +func (a *Session) Agent() *agent.Agent { return a.r } +func (a *Session) Bus() *pubsub.Bus { return a.bus } +func (a *Session) ToolCallCount() int64 { return a.r.ToolCallCount() } +func (a *Session) CtxSz() string { return a.r.CtxSz() } +func (a *Session) Cost() string { return a.r.CostStr() } +func (a *Session) Usage() string { return a.r.UsageStr() } +func (a *Session) Context() []backend.Message { return a.r.Context() } +func (a *Session) SystemPrompt() string { return a.r.SystemPrompt() } +func (a *Session) CompactionModel() string { return a.r.CompactionModel() } +func (a *Session) SetCompactionModel(m string) { a.r.SetCompactionModel(m) } + +func (a *Session) GenerationParams() backend.GenerationParams { + return a.r.GenParams() +} + +func (a *Session) SetGenerationParams(params backend.GenerationParams) error { + if a.IsRunning() { + return fmt.Errorf("cannot change params while agent is running") } + a.r.SetGenParams(params) return nil } -func (a *Session) SignalDetached(pid, signal int) error { - if srv := a.execServer(); srv != nil { - type signaler interface { - SignalDetached(int, syscall.Signal) error - } - if sg, ok := srv.(signaler); ok { - return sg.SignalDetached(pid, syscall.Signal(signal)) - } - } - return fmt.Errorf("no execute server available") +func (a *Session) ListModels() string { + models := a.r.ListModels() + return strings.Join(models, "\n") } -func (a *Session) GetDetachedOutput(pid int) (string, error) { - if srv := a.execServer(); srv != nil { - type outputGetter interface { - GetDetachedOutput(int) (string, error) - } - if og, ok := srv.(outputGetter); ok { - return og.GetDetachedOutput(pid) - } - } - return "", fmt.Errorf("no execute server available") +func (a *Session) Interrupt(cause error) bool { + a.log.Debug("Interrupt() cause=%v", cause) + return a.r.Interrupt(cause) } -func (a *Session) DismissDetached(pid int) bool { - if srv := a.execServer(); srv != nil { - type dismisser interface { - DismissDetached(int) bool - } - if d, ok := srv.(dismisser); ok { - return d.DismissDetached(pid) - } - } - return false +func (a *Session) Queue(prompt string) { + a.r.Queue(prompt) + a.bus.Publish("queued", prompt) } +func (a *Session) PopQueue() (string, bool) { return a.r.PopQueue() } + func (a *Session) InjectSystemEvent(content string) { a.Queue("\n" + content + "\n") } -// 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 - } -} - -func (a *Session) Reactions() map[string]string { - result := make(map[string]string) - if a.r.history == nil { - return result - } - for _, reaction := range a.r.history.Reactions { - result[reaction.ResponseID] = reaction.Emoji - } - return result -} - -func (a *Session) React(emoji string) { - _ = a.ReactTo("", emoji) -} - +func (a *Session) Reactions() map[string]string { return a.r.Reactions() } +func (a *Session) React(emoji string) { _ = a.r.React("", emoji) } func (a *Session) ReactTo(responseID, emoji string) error { - if a.r.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(a.r.history.messages) - 1; i >= 0; i-- { - if a.r.history.messages[i].Role == "assistant" { - responseID = a.r.history.messages[i].ID - break - } - } - } - if responseID == "" { - return fmt.Errorf("no assistant response to react to") - } - found := false - for i := range a.r.history.messages { - if a.r.history.messages[i].Role == "assistant" && a.r.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 a.r.history.Reactions { - if a.r.history.Reactions[i].ResponseID == responseID { - if a.r.history.Reactions[i].Emoji == emoji { - return nil - } - a.r.history.Reactions[i] = reaction - replaced = true - break - } - } - if !replaced { - a.r.history.Reactions = append(a.r.history.Reactions, reaction) - } - a.r.history.recomputeReactionCounts() - a.saveSession() - return nil -} - -func (a *Session) AgentName() string { - v := a.r.Name() - a.log.Debug("AgentName() = %q", v) - return v -} -func (a *Session) BackendName() string { - v := a.r.BackendName() - a.log.Debug("BackendName() = %q", v) - return v -} -func (a *Session) ModelName() string { - v := a.r.ModelName() - a.log.Debug("ModelName() = %q", v) - return v -} - -// Agent returns the active agent. -func (a *Session) Agent() *Agent { return a.r } - -func (a *Session) State() string { - return a.r.State() -} - -func (a *Session) setState(state string) { - a.log.Debug("state -> %q", state) - a.r.SetState(state) -} - -// WaitChange blocks until the named field changes from current, then returns -// the new value. Returns ("", false) if ctx is cancelled. -func (a *Session) WaitChange(ctx context.Context, field, current string) (string, bool) { - // For state changes, delegate directly to the agent. - if field == WatchState { - return a.r.WaitChange(ctx, field, current) - } - - // Other watched fields (usage, ctxsz, cwd, agent) are session-level; - // they change as a side effect of agent activity, so we still listen - // on the agent's change signal. - read := func() string { - switch field { - case WatchUsage: - return a.Usage() - case WatchCtxSz: - return a.CtxSz() - case WatchCWD: - return a.CWD() - case WatchAgent: - return a.AgentName() - } - return "" - } - - stop := context.AfterFunc(ctx, func() { - a.r.changeMu.Lock() - a.r.changeCond.Broadcast() - a.r.changeMu.Unlock() - }) - defer stop() - - a.r.changeMu.Lock() - defer a.r.changeMu.Unlock() - for ctx.Err() == nil { - if v := read(); v != current { - return v, true - } - a.r.changeCond.Wait() - } - return "", false -} - -func (a *Session) Reply() string { - r := a.r.Reply() - a.log.Debug("Reply() len=%d", len(r)) - return r + return a.r.React(responseID, emoji) } // CWD returns the current working directory for tool execution. func (a *Session) CWD() string { - cwd := a.r.Cwd() - if cwd != "" { - a.log.Debug("CWD() = %q", cwd) - return cwd + if c := a.r.Cwd(); c != "" { + return c } wd, _ := os.Getwd() - a.log.Debug("CWD() = %q (from getwd)", wd) return wd } -// SetCWD changes the working directory for tool execution and updates the -// system prompt. Returns an error if the path does not exist. +// SetCWD validates and sets the working directory. func (a *Session) SetCWD(dir string) error { - a.log.Debug("SetCWD(%q)", dir) dir = paths.ExpandHome(dir) if dir != "" { if _, err := os.Stat(dir); err != nil { return fmt.Errorf("cwd: %w", err) } } - oldCwd := a.r.Cwd() - a.r.SetCwd(dir) - // Update cwd references in the system prompt. - if oldCwd != "" && dir != "" && oldCwd != dir { - a.r.runtime.Preamble = strings.ReplaceAll(a.r.runtime.Preamble, oldCwd, dir) - } - // Propagate to any tool server that knows how to handle it (e.g. execute). - if a.r.runtime != nil && a.r.runtime.Dispatcher != nil { - if srv, ok := a.r.runtime.Dispatcher.GetServer("execute"); ok { - if ws, ok := srv.(tools.CWDSetter); ok { - ws.SetCWD(dir) - } - } - } - a.r.notifyChange() + a.r.SetCWD(dir) return nil } -// SetSessionID renames the session. It updates the in-memory ID, renames -// persisted files on disk, and propagates to the execute server env. +// SetSessionID renames the session. func (a *Session) SetSessionID(newID string) error { - a.log.Debug("SetSessionID(%q) old=%q", newID, a.id) oldID := a.id if oldID == newID { return nil } - // Rename active persisted files on disk. if a.sessionsDir != "" && oldID != "" { for _, suffix := range []string{".json", ".compaction.jsonl"} { oldPath := a.activeSessionPath(oldID, suffix) @@ -822,23 +291,49 @@ func (a *Session) SetSessionID(newID string) error { } } a.id = newID - // Update session ID references in the system prompt. - a.r.runtime.Preamble = strings.ReplaceAll(a.r.runtime.Preamble, oldID, newID) - // Rename tmpdir so isread markers remain valid after rename. + a.r.RenamePreamble(oldID, newID) oldTemp := filepath.Join(ollieTmpDir(), oldID) newTemp := filepath.Join(ollieTmpDir(), newID) if _, err := os.Stat(oldTemp); err == nil { os.Rename(oldTemp, newTemp) //nolint:errcheck } - a.pushSessionEnv() + a.r.SetSessionEnv(newID) return nil } -// defaultContextLength is used when the backend cannot report the model's -// actual context window (e.g. CodeWhisperer). 128k tokens is a safe default -// for modern models. -const defaultContextLength = 128000 -const defaultToolResultMaxBytes = 131072 +// WaitChange blocks until the named field changes from current. +func (a *Session) WaitChange(ctx context.Context, field, current string) (string, bool) { + if field == agent.WatchState { + return a.r.WaitChange(ctx, field, current) + } + // Other fields — read via session methods, use agent's change signal. + read := func() string { + switch field { + case "usage": + return a.Usage() + case "ctxsz": + return a.CtxSz() + case "cwd": + return a.CWD() + case "agent": + return a.AgentName() + } + return "" + } + stop := context.AfterFunc(ctx, func() { a.r.BroadcastChange() }) + defer stop() + for ctx.Err() == nil { + if v := read(); v != current { + return v, true + } + a.r.WaitForChange(ctx) + } + return "", false +} + +func (a *Session) emit(ev agent.Event) { + a.bus.Publish("event", ev) +} func (a *Session) activeSessionPath(id, suffix string) string { return filepath.Join(a.sessionsDir, "active", id+suffix) @@ -853,7 +348,6 @@ func (a *Session) saveSession() { a.saveMu.Unlock() } -// flushSave immediately persists the session if dirty. func (a *Session) flushSave() { a.saveMu.Lock() dirty := a.saveDirty @@ -863,10 +357,7 @@ func (a *Session) flushSave() { a.saveTimer = nil } a.saveMu.Unlock() - if !dirty { - return - } - if a.r.history == nil || a.id == "" || a.sessionsDir == "" { + if !dirty || a.id == "" || a.sessionsDir == "" { return } path := a.activeSessionPath(a.id, ".json") @@ -874,288 +365,87 @@ func (a *Session) flushSave() { a.log.Error("session save: %v", err) return } - if err := a.r.history.saveToFull(path, a.id, a.r.agentName, - a.r.runtime.Backend.Name(), a.r.runtime.Backend.Model(), a.CWD(), a.remote); err != nil { + if err := a.r.SaveFull(path, a.id, a.CWD(), a.remote); err != nil { a.log.Error("session save: %v", err) } } -// SaveSession writes the current session state to the given path, including -// backend and model metadata for external restore. +// SaveSession writes the current session state to the given path. func (a *Session) SaveSession(path string) error { - a.mu.RLock() - defer a.mu.RUnlock() - if a.r.history == nil { - return fmt.Errorf("no active session") - } - return a.r.history.saveToFull(path, a.id, a.r.agentName, - a.r.runtime.Backend.Name(), a.r.runtime.Backend.Model(), a.CWD(), a.remote) + return a.r.SaveFull(path, a.id, a.CWD(), a.remote) } -// Interrupt cancels the current in-progress agent turn. -// Returns true if an action was running and was cancelled. -func (a *Session) Interrupt(cause error) bool { - a.log.Debug("Interrupt() cause=%v", cause) - return a.r.Interrupt(cause) -} - -func (a *Session) Inject(prompt string) { - // If an inject is already pending, fall back to the normal FIFO so nothing - // is lost. Use CompareAndSwap to avoid a race between the nil check and store. - if !a.r.pendingInject.CompareAndSwap(nil, &prompt) { - a.r.fifo.Push(prompt) - return - } - a.emit(Event{Role: "info", Content: "\n"}) - a.emit(Event{Role: "user", Content: prompt}) -} - -func (a *Session) injectRewrite(prompt string) { - a.r.pendingInject.Store(&prompt) - a.emit(Event{Role: "info", Content: "\n"}) - a.emit(Event{Role: "user", Content: prompt}) -} - -func (a *Session) Queue(prompt string) { - a.r.fifo.Push(prompt) - a.bus.Publish("queued", prompt) -} - -func (a *Session) drainQueue() { - if prompt, ok := a.r.fifo.Pop(); ok { - a.Submit(context.Background(), prompt) - } -} - -func (a *Session) Bus() *pubsub.Bus { - return a.bus -} - -func (a *Session) emit(ev Event) { - a.bus.Publish("event", ev) -} - -func (a *Session) PopQueue() (string, bool) { - return a.r.fifo.Pop() -} - -func (a *Session) IsRunning() bool { - v := a.r.IsRunning() - a.log.Debug("IsRunning() = %v", v) - return v -} - -func (a *Session) CtxSz() string { - if a.r.history == nil { - a.log.Debug("CtxSz() no session") - return "no active session" - } - ctxLen := a.r.runtime.Backend.ContextLength(context.Background()) - if ctxLen <= 0 { - ctxLen = defaultContextLength - } - estimated := a.r.history.estimateTokens() - pct := estimated * 100 / ctxLen - v := fmt.Sprintf("%d / %d (%d%%)", estimated, ctxLen, pct) - a.log.Debug("CtxSz() = %q", v) - return v -} - -func (a *Session) Cost() string { - if a.r.history == nil { - return "no active session" - } - return fmt.Sprintf("costLast=$%.4f\ncostSession=$%.4f\n", - a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD) -} - -func (a *Session) Usage() string { - if a.r.history == nil { - a.log.Debug("Usage() no session") - return "no active session" - } - str := fmt.Sprintf("%d in, %d out, %d requests", - a.r.history.TotalInputTokens, a.r.history.TotalOutputTokens, - a.r.history.TotalRequests) - if a.r.history.TotalCachedInputTokens > 0 { - str += fmt.Sprintf(", %d cached", a.r.history.TotalCachedInputTokens) - } - if a.r.history.Estimated { - str += " [estimated]" - } - a.log.Debug("Usage() = %q", str) - return str -} - -func (a *Session) Context() []backend.Message { - a.mu.RLock() - var msgs []backend.Message - if a.r.history != nil { - msgs = slices.Clone(a.r.history.history()) - } - a.mu.RUnlock() - if a.r.runtime.Preamble != "" { - msgs = append([]backend.Message{{Role: "system", Content: a.r.runtime.Preamble}}, msgs...) - } - return msgs -} - -func (a *Session) SystemPrompt() string { - a.log.Debug("SystemPrompt() len=%d", len(a.r.runtime.Preamble)) - return a.r.runtime.Preamble -} - -func (a *Session) GenerationParams() backend.GenerationParams { - a.mu.RLock() - defer a.mu.RUnlock() - return a.r.runtime.GenParams -} - -func (a *Session) CompactionModel() string { - a.mu.RLock() - defer a.mu.RUnlock() - return a.r.runtime.CompactionModel -} - -func (a *Session) SetCompactionModel(model string) { - a.mu.Lock() - defer a.mu.Unlock() - a.r.runtime.CompactionModel = model -} - -func (a *Session) SetGenerationParams(params backend.GenerationParams) error { - if a.IsRunning() { - return fmt.Errorf("cannot change params while agent is running") - } - a.mu.Lock() - a.r.runtime.GenParams = params - a.mu.Unlock() - return nil -} - -func (a *Session) ListModels() string { - a.log.Debug("ListModels()") - models := a.r.runtime.Backend.Models(context.Background()) - slices.Sort(models) - return strings.Join(models, "\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 -// is found. -func firstSentence(s string) string { - for i, r := range s { - if r == '.' || r == '\n' { - return strings.TrimSpace(s[:i+1]) +// Detach operations delegate to the agent's execute server. +func (a *Session) Detach() bool { + if srv := a.r.ExecServer(); srv != nil { + if d, ok := srv.(interface{ Detach() bool }); ok { + return d.Detach() } } - if len(s) > 80 { - return s[:77] + "..." - } - return s + return false } -// Submit processes one line of user input: slash commands are dispatched -// immediately; any other input is delegated to the agent for turn execution. -func (a *Session) Submit(ctx context.Context, input string) { - a.log.Debug("Submit() input_len=%d running=%v", len(input), a.IsRunning()) - if input == "" { - return - } - - // Commands are handled at the session level (they reference sessionsDir, - // listHandlers, etc.). If the input is a command, handle it and return. - if a.IsRunning() { - if a.handleCommand(ctx, input) { - return - } - } else { - // Serialize commands with turns — acquire submitMu to check. - a.r.submitMu.Lock() - if a.handleCommand(ctx, input) { - a.r.submitMu.Unlock() - return - } - a.r.submitMu.Unlock() - } - - // Delegate to the agent for turn execution (or FIFO queueing if running). - a.r.Submit(ctx, input) +// DetachedInfo describes a detached process. +type DetachedInfo struct { + PID int + Command string + Started int64 + Exited bool + ExitCode int } -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, +func (a *Session) ListDetached() []DetachedInfo { + srv := a.r.ExecServer() + if srv == nil { + return nil + } + type listDetacher interface{ ListDetachedRaw() []any } + ld, ok := srv.(listDetacher) + if !ok { + return nil + } + raw := ld.ListDetachedRaw() + out := make([]DetachedInfo, 0, len(raw)) + for _, r := range raw { + if m, ok := r.(map[string]any); ok { + di := DetachedInfo{} + if v, ok := m["pid"].(int); ok { di.PID = v } + if v, ok := m["command"].(string); ok { di.Command = v } + if v, ok := m["started"].(int64); ok { di.Started = v } + if v, ok := m["exited"].(bool); ok { di.Exited = v } + if v, ok := m["exit_code"].(int); ok { di.ExitCode = v } + out = append(out, di) } } 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, - }, - }) +func (a *Session) SignalDetached(pid, signal int) error { + if srv := a.r.ExecServer(); srv != nil { + type signaler interface{ SignalDetached(int, syscall.Signal) error } + if sg, ok := srv.(signaler); ok { + return sg.SignalDetached(pid, syscall.Signal(signal)) } } - return strings.Join(parts, "\n"), contentBlocks, result.IsError + return fmt.Errorf("no execute server available") } -// WatchField names supported by Session.WaitChange. -const ( - WatchState = "state" - WatchUsage = "usage" - WatchCtxSz = "ctxsz" - WatchCWD = "cwd" - WatchAgent = "agent" -) - -// ErrInterrupted is returned when the user cancels an agent turn (Ctrl-C). -var ErrInterrupted = errors.New("interrupted") - -// 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 +func (a *Session) GetDetachedOutput(pid int) (string, error) { + if srv := a.r.ExecServer(); srv != nil { + type outputGetter interface{ GetDetachedOutput(int) (string, error) } + if og, ok := srv.(outputGetter); ok { + return og.GetDetachedOutput(pid) + } + } + return "", fmt.Errorf("no execute server available") } -// EventHandler receives events from the agent. -type EventHandler func(Event) - -// DetachedInfo describes a detached process for external consumers. -type DetachedInfo struct { - PID int - Command string - Started int64 // unix timestamp - Exited bool - ExitCode int +func (a *Session) DismissDetached(pid int) bool { + if srv := a.r.ExecServer(); srv != nil { + type dismisser interface{ DismissDetached(int) bool } + if d, ok := srv.(dismisser); ok { + return d.DismissDetached(pid) + } + } + return false } diff --git a/session/waitchange_test.go b/session/waitchange_test.go deleted file mode 100644 index 3fc5aa3..0000000 --- a/session/waitchange_test.go +++ /dev/null @@ -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) - } - } -}