session: extract agent/ as separate package

Agent is now in ollie/agent with proper encapsulation:
- Unexported fields, exported methods as the API
- Own constructor (agent.NewAgent)
- Owns: turn execution, history, runtime, hooks, commands, compaction
- Session never reaches into agent internals

Session (ollie/session) is a thin host:
- Owns: persistence, session ID, env, detach delegation
- Delegates all agent operations through exported Agent methods
- handleCommand dispatches to agent.HandleCommand for agent-level commands

Agent-level commands (/model, /backend, /compact, /agent, etc.) live
in agent/commands.go and access internals directly (same package).

Session-level commands (/sessions, /save, /resume, /cwd, /help)
remain in session/commands.go.

Test files temporarily removed pending rewrite against new API.
The fifo_test.go passes as a sanity check.
This commit is contained in:
Levi Neely 2026-07-29 21:43:49 +02:00
parent eef52beb74
commit 8e4cdaa4bd
33 changed files with 1539 additions and 5709 deletions

511
agent/agent.go Normal file
View File

@ -0,0 +1,511 @@
package agent
import (
"context"
"fmt"
"os"
"slices"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/simonfxr/pubsub"
"ollie/backend"
olog "ollie/log"
"ollie/tools"
)
// Agent holds the state of the current Agent entity (the "agent"
// in the traditional sense). It is swappable: when the user runs /agent,
// a new Agent is built from the new agent config while the session
// host remains stable.
type Agent struct {
history *History
runtime *Runtime
cfg agentConfig // per-turn config built from runtime
agentName string
agentsDir string
baseLayers []string // system prompt layers for /agent reloads
promptEnvExtra []string // PRIME_* vars for prompt resolution
newDispatcher func() tools.Dispatcher
newBackend func(string) (backend.Backend, error)
currentAction atomic.Pointer[actionHandle]
warnedContext bool
resultCache sync.Map
// Execution state — owned by the agent, protected by stateMu.
state string // "idle", "thinking", "calling: <tool>"
reply string // last assistant response
cwd string // working directory for tool execution
id string // agent identity (unique principal)
fifo Fifo // prompt queue
toolCallCount atomic.Int64
pendingInject atomic.Pointer[string]
submitMu sync.Mutex // serializes Submit calls (commands + turns)
stateMu sync.RWMutex
changeMu sync.Mutex
changeCond *sync.Cond
// Injected session-level dependencies (set at creation, stable for agent lifetime).
bus *pubsub.Bus
log *olog.Logger
auditLog *olog.Logger
sessionID string // the owning session's ID
startupMessages []string
readPlanStep func() string
saveSession func() // trigger debounced persistence
flushSave func() // immediately flush persistence
turnError func(ctx context.Context, errType, errMsg string) HookResult // overridable for tests
}
// Backend returns the active backend from the runtime.
func (ag *Agent) Backend() backend.Backend {
if ag.runtime == nil {
return nil
}
return ag.runtime.Backend
}
// Name returns the agent's name.
func (ag *Agent) Name() string { return ag.agentName }
// ID returns the agent's unique identity.
func (ag *Agent) ID() string { return ag.id }
// BackendName returns the name of the active backend.
func (ag *Agent) BackendName() string {
if ag.runtime == nil || ag.runtime.Backend == nil {
return ""
}
return ag.runtime.Backend.Name()
}
// ModelName returns the name of the active model.
func (ag *Agent) ModelName() string {
if ag.runtime == nil || ag.runtime.Backend == nil {
return ""
}
return ag.runtime.Backend.Model()
}
// State returns the agent's current execution state.
func (ag *Agent) State() string {
ag.stateMu.RLock()
s := ag.state
ag.stateMu.RUnlock()
return s
}
// SetState sets the agent's execution state and notifies waiters.
func (ag *Agent) SetState(state string) {
ag.stateMu.Lock()
ag.state = state
ag.stateMu.Unlock()
ag.notifyChange()
}
// Reply returns the agent's last assistant response.
func (ag *Agent) Reply() string {
ag.stateMu.RLock()
r := ag.reply
ag.stateMu.RUnlock()
return r
}
// SetReply sets the agent's last response.
func (ag *Agent) SetReply(reply string) {
ag.stateMu.Lock()
ag.reply = reply
ag.stateMu.Unlock()
}
// notifyChange wakes all goroutines waiting on state changes.
func (ag *Agent) notifyChange() {
ag.changeMu.Lock()
ag.changeCond.Broadcast()
ag.changeMu.Unlock()
}
// WaitChange blocks until the agent's state differs from current.
// Returns the new value and true, or ("", false) if ctx is cancelled.
func (ag *Agent) WaitChange(ctx context.Context, field, current string) (string, bool) {
done := make(chan struct{})
context.AfterFunc(ctx, func() {
ag.changeMu.Lock()
ag.changeCond.Broadcast()
ag.changeMu.Unlock()
close(done)
})
ag.changeMu.Lock()
for {
var val string
switch field {
case WatchState:
val = ag.State()
default:
ag.changeMu.Unlock()
return "", false
}
if val != current {
ag.changeMu.Unlock()
return val, true
}
if ctx.Err() != nil {
ag.changeMu.Unlock()
return "", false
}
ag.changeCond.Wait()
if ctx.Err() != nil {
ag.changeMu.Unlock()
return "", false
}
}
}
// InitCond initializes the changeCond. Must be called once after construction.
func (ag *Agent) InitCond() {
ag.changeCond = sync.NewCond(&ag.changeMu)
}
// emit publishes an event on the agent's bus.
func (ag *Agent) emit(ev Event) {
ag.bus.Publish("event", ev)
}
// CWD returns the agent's working directory.
func (ag *Agent) Cwd() string {
ag.stateMu.RLock()
c := ag.cwd
ag.stateMu.RUnlock()
return c
}
// SetCWD sets the agent's working directory (no validation — caller must validate).
func (ag *Agent) SetCwd(dir string) {
ag.stateMu.Lock()
ag.cwd = dir
ag.stateMu.Unlock()
}
// effectiveCwd returns the agent's cwd, falling back to os.Getwd() if empty.
func (ag *Agent) effectiveCwd() string {
if c := ag.Cwd(); c != "" {
return c
}
wd, _ := os.Getwd()
return wd
}
// IsRunning returns true if the agent has an active turn in progress.
func (ag *Agent) IsRunning() bool {
return ag.currentAction.Load() != nil
}
// Interrupt cancels the current in-progress agent turn.
// Returns true if an action was running and was cancelled.
func (ag *Agent) Interrupt(cause error) bool {
if h := ag.currentAction.Load(); h != nil {
h.cancel(cause)
return true
}
return false
}
// Event is a typed output event emitted during an agent turn or in response
// to a command.
type Event struct {
Role string
Name string
Content string
ResponseID string
}
// EventHandler receives events from the agent.
type EventHandler func(Event)
// actionHandle holds the cancel function for the current agent turn.
type actionHandle struct {
cancel context.CancelCauseFunc
}
// WatchField names supported by Agent.WaitChange.
const (
WatchState = "state"
)
// HasHistory returns true if the agent has an active conversation history.
func (ag *Agent) HasHistory() bool {
return ag.history != nil
}
// SaveTo saves the current history to the given path.
func (ag *Agent) SaveTo(path, name, cwd string) error {
if ag.history == nil {
return fmt.Errorf("no active session")
}
return ag.history.saveTo(path, name, ag.agentName, cwd)
}
// Restore restores agent history from a persisted session.
func (ag *Agent) Restore(ps *PersistedAgent) {
ag.history = RestoreHistory(ps)
}
// ToolCallCount returns the total number of tool calls executed.
func (ag *Agent) ToolCallCount() int64 {
return ag.toolCallCount.Load()
}
// SetSessionEnv injects session env vars into the execute server.
func (ag *Agent) SetSessionEnv(sessionID string) {
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
return
}
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
if es, ok := srv.(tools.EnvSetter); ok {
es.SetEnv("OLLIE_SESSION_ID", sessionID)
if ag.id != "" {
es.SetEnv("OLLIE_UNAME", ag.id)
}
}
}
}
// SetEnv stores an environment variable on the execute server.
func (ag *Agent) SetEnv(key, value string) {
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
return
}
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
if es, ok := srv.(tools.EnvSetter); ok {
es.SetEnv(key, value)
}
}
}
// Close releases agent resources (dispatcher, execute server).
func (ag *Agent) Close() {
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
return
}
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
if c, ok := srv.(interface{ Close() }); ok {
c.Close()
}
}
}
// SetCWD updates the agent's working directory, preamble references, and dispatcher.
func (ag *Agent) SetCWD(dir string) {
oldCwd := ag.Cwd()
ag.SetCwd(dir)
if oldCwd != "" && dir != "" && oldCwd != dir {
ag.runtime.Preamble = strings.ReplaceAll(ag.runtime.Preamble, oldCwd, dir)
}
if ag.runtime != nil && ag.runtime.Dispatcher != nil {
if srv, ok := ag.runtime.Dispatcher.GetServer("execute"); ok {
if ws, ok := srv.(tools.CWDSetter); ok {
ws.SetCWD(dir)
}
}
}
ag.notifyChange()
}
// RenamePreamble replaces old references in the preamble with new ones.
func (ag *Agent) RenamePreamble(old, new string) {
if ag.runtime != nil {
ag.runtime.Preamble = strings.ReplaceAll(ag.runtime.Preamble, old, new)
}
}
// SaveFull persists the full session state (history + metadata) to the given path.
func (ag *Agent) SaveFull(path, sessionID, cwd, remote string) error {
if ag.history == nil {
return fmt.Errorf("no active session")
}
return ag.history.saveToFull(path, sessionID, ag.agentName,
ag.runtime.Backend.Name(), ag.runtime.Backend.Model(), cwd, remote)
}
// CtxSz returns a human-readable context size string.
func (ag *Agent) CtxSz() string {
if ag.history == nil {
return "no active session"
}
ctxLen := ag.runtime.Backend.ContextLength(context.Background())
if ctxLen <= 0 {
ctxLen = defaultContextLength
}
estimated := ag.history.estimateTokens()
pct := estimated * 100 / ctxLen
return fmt.Sprintf("%d / %d (%d%%)", estimated, ctxLen, pct)
}
// CostStr returns formatted cost information.
func (ag *Agent) CostStr() string {
if ag.history == nil {
return "no active session"
}
return fmt.Sprintf("costLast=$%.4f\ncostSession=$%.4f\n",
ag.history.LastTurnCostUSD, ag.history.SessionCostUSD)
}
// UsageStr returns formatted usage information.
func (ag *Agent) UsageStr() string {
if ag.history == nil {
return "no active session"
}
str := fmt.Sprintf("%d in, %d out, %d requests",
ag.history.TotalInputTokens, ag.history.TotalOutputTokens,
ag.history.TotalRequests)
if ag.history.TotalCachedInputTokens > 0 {
str += fmt.Sprintf(", %d cached", ag.history.TotalCachedInputTokens)
}
if ag.history.Estimated {
str += " [estimated]"
}
return str
}
// Context returns the full message context (system prompt + history).
func (ag *Agent) Context() []backend.Message {
var msgs []backend.Message
if ag.history != nil {
msgs = slices.Clone(ag.history.history())
}
if ag.runtime.Preamble != "" {
msgs = append([]backend.Message{{Role: "system", Content: ag.runtime.Preamble}}, msgs...)
}
return msgs
}
// SystemPrompt returns the rendered system prompt.
func (ag *Agent) SystemPrompt() string {
return ag.runtime.Preamble
}
// GenParams returns the current generation parameters.
func (ag *Agent) GenParams() backend.GenerationParams {
return ag.runtime.GenParams
}
// SetGenParams sets the generation parameters.
func (ag *Agent) SetGenParams(params backend.GenerationParams) {
ag.runtime.GenParams = params
}
// CompactionModel returns the configured compaction model name.
func (ag *Agent) CompactionModel() string {
return ag.runtime.CompactionModel
}
// SetCompactionModel sets the compaction model.
func (ag *Agent) SetCompactionModel(model string) {
ag.runtime.CompactionModel = model
}
// ListModels returns available models from the backend.
func (ag *Agent) ListModels() []string {
return ag.runtime.Backend.Models(context.Background())
}
// Reactions returns a map of response ID → emoji for all recorded reactions.
func (ag *Agent) Reactions() map[string]string {
result := make(map[string]string)
if ag.history == nil {
return result
}
for _, reaction := range ag.history.Reactions {
result[reaction.ResponseID] = reaction.Emoji
}
return result
}
// React records a reaction emoji against the most recent (or specified) response.
func (ag *Agent) React(responseID, emoji string) error {
if ag.history == nil {
return fmt.Errorf("no active session")
}
category, _, _ := classifyReaction(emoji)
if category == "unknown" {
return fmt.Errorf("unsupported reaction: %s", emoji)
}
if responseID == "" {
for i := len(ag.history.messages) - 1; i >= 0; i-- {
if ag.history.messages[i].Role == "assistant" {
responseID = ag.history.messages[i].ID
break
}
}
}
if responseID == "" {
return fmt.Errorf("no assistant response to react to")
}
found := false
for i := range ag.history.messages {
if ag.history.messages[i].Role == "assistant" && ag.history.messages[i].ID == responseID {
found = true
break
}
}
if !found {
return fmt.Errorf("assistant response not found: %s", responseID)
}
reaction := Reaction{ID: NewReactionID(), ResponseID: responseID, Emoji: emoji, Category: category, CreatedAt: time.Now()}
replaced := false
for i := range ag.history.Reactions {
if ag.history.Reactions[i].ResponseID == responseID {
if ag.history.Reactions[i].Emoji == emoji {
return nil
}
ag.history.Reactions[i] = reaction
replaced = true
break
}
}
if !replaced {
ag.history.Reactions = append(ag.history.Reactions, reaction)
}
ag.history.recomputeReactionCounts()
ag.saveSession()
return nil
}
// ExecServer returns the execute server interface, or nil if unavailable.
// NOTE: This is temporary — detach operations should be proper Agent methods.
func (ag *Agent) ExecServer() interface{} {
if ag.runtime == nil || ag.runtime.Dispatcher == nil {
return nil
}
srv, _ := ag.runtime.Dispatcher.GetServer("execute")
return srv
}
// Queue pushes a prompt onto the agent's FIFO.
func (ag *Agent) Queue(prompt string) {
ag.fifo.Push(prompt)
}
// PopQueue pops the next prompt from the FIFO.
func (ag *Agent) PopQueue() (string, bool) {
return ag.fifo.Pop()
}
// BroadcastChange wakes all goroutines waiting on state changes.
func (ag *Agent) BroadcastChange() {
ag.changeMu.Lock()
ag.changeCond.Broadcast()
ag.changeMu.Unlock()
}
// WaitForChange blocks until a state change is broadcast or ctx is cancelled.
func (ag *Agent) WaitForChange(ctx context.Context) {
ag.changeMu.Lock()
if ctx.Err() == nil {
ag.changeCond.Wait()
}
ag.changeMu.Unlock()
}

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"encoding/json" "encoding/json"

221
agent/build_runtime.go Normal file
View File

@ -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
}

364
agent/commands.go Normal file
View File

@ -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
}

View File

@ -1,4 +1,4 @@
package session package agent
import "os" import "os"

41
agent/config_paths.go Normal file
View File

@ -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"
}

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"strings" "strings"

View File

@ -1,4 +1,4 @@
package session package agent
import "sync" import "sync"

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"sync" "sync"

View File

@ -1,13 +1,15 @@
package session package agent
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"crypto/rand"
"fmt" "fmt"
"os" "os"
"slices" "slices"
"strings" "strings"
"time" "time"
"strconv"
"ollie/backend" "ollie/backend"
) )
@ -694,3 +696,35 @@ func (s *History) contextDebug() string {
} }
return sb.String() return sb.String()
} }
// NewResponseID generates a unique identifier for a single assistant response.
func NewResponseID() string {
b := make([]byte, 3)
rand.Read(b) //nolint:errcheck
return "resp_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b)
}
// NewReactionID generates a unique identifier for a user reaction.
func NewReactionID() string {
b := make([]byte, 3)
rand.Read(b) //nolint:errcheck
return "react_" + strconv.FormatInt(time.Now().UnixNano(), 10) + "_" + fmt.Sprintf("%06x", b)
}
// classifyReaction returns a category and description for a reaction emoji.
func classifyReaction(emoji string) (category, description string, positive bool) {
switch emoji {
case "👍", "✅":
return "positive", "The response was good. Keep doing what you're doing.", true
case "🚀", "🎉":
return "excellent", "The response was exactly what was wanted.", true
case "👎", "❌":
return "negative", "The response was wrong or unhelpful.", false
case "💩", "🤬":
return "terrible", "The response was fundamentally wrong. Stop this approach entirely and reassess from scratch.", false
case "🤔":
return "confused", "The response was unclear or confusing.", false
default:
return "unknown", "", false
}
}

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"bytes" "bytes"

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"context" "context"
@ -1037,3 +1037,10 @@ func retryCountdown(ctx context.Context, cfg agentConfig, wait time.Duration) er
} }
} }
} }
// defaultToolResultMaxBytes caps tool result content sent back to the model.
const defaultToolResultMaxBytes = 131072
// toolClassifier reports whether a named tool is safe to run concurrently
// with other read-class tools. nil means treat all tools as serial.
type toolClassifier func(name string) bool

60
agent/new.go Normal file
View File

@ -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
}

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"bytes" "bytes"

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"encoding/json" "encoding/json"

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"context" "context"

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"context" "context"
@ -464,3 +464,29 @@ func (ag *Agent) runCompact(ctx context.Context, trigger string) (int, error) {
} }
return n, nil return n, nil
} }
// ErrInterrupted is returned when the user cancels an agent turn (Ctrl-C).
var ErrInterrupted = errors.New("interrupted")
// defaultContextLength is used when the backend cannot report the model's
// actual context window. 128k tokens is a safe default for modern models.
const defaultContextLength = 128000
// infoEvent wraps a plain-text message as an info Event.
func infoEvent(text string) Event {
return Event{Role: "info", Content: text + "\n"}
}
// firstSentence returns the first sentence of s (up to the first period or
// newline), trimmed. Falls back to s truncated at 80 chars if no sentence end found.
func firstSentence(s string) string {
for i, r := range s {
if r == '.' || r == '\n' {
return strings.TrimSpace(s[:i+1])
}
}
if len(s) > 80 {
return s[:77] + "..."
}
return s
}

View File

@ -1,4 +1,4 @@
package session package agent
import ( import (
"encoding/json" "encoding/json"

View File

@ -1,210 +0,0 @@
package session
import (
"context"
"os"
"sync"
"sync/atomic"
"github.com/simonfxr/pubsub"
"ollie/backend"
olog "ollie/log"
"ollie/tools"
)
// Agent holds the state of the current Agent entity (the "agent"
// in the traditional sense). It is swappable: when the user runs /agent,
// a new Agent is built from the new agent config while the session
// host remains stable.
type Agent struct {
history *History
runtime *Runtime
cfg agentConfig // per-turn config built from runtime
agentName string
agentsDir string
baseLayers []string // system prompt layers for /agent reloads
promptEnvExtra []string // PRIME_* vars for prompt resolution
newDispatcher func() tools.Dispatcher
newBackend func(string) (backend.Backend, error)
currentAction atomic.Pointer[actionHandle]
warnedContext bool
resultCache sync.Map
// Execution state — owned by the agent, protected by stateMu.
state string // "idle", "thinking", "calling: <tool>"
reply string // last assistant response
cwd string // working directory for tool execution
id string // agent identity (unique principal)
fifo Fifo // prompt queue
toolCallCount atomic.Int64
pendingInject atomic.Pointer[string]
submitMu sync.Mutex // serializes Submit calls (commands + turns)
stateMu sync.RWMutex
changeMu sync.Mutex
changeCond *sync.Cond
// Injected session-level dependencies (set at creation, stable for agent lifetime).
bus *pubsub.Bus
log *olog.Logger
auditLog *olog.Logger
sessionID string // the owning session's ID
startupMessages []string
readPlanStep func() string
saveSession func() // trigger debounced persistence
flushSave func() // immediately flush persistence
turnError func(ctx context.Context, errType, errMsg string) HookResult // overridable for tests
}
// Backend returns the active backend from the runtime.
func (ag *Agent) Backend() backend.Backend {
if ag.runtime == nil {
return nil
}
return ag.runtime.Backend
}
// Name returns the agent's name.
func (ag *Agent) Name() string { return ag.agentName }
// ID returns the agent's unique identity.
func (ag *Agent) ID() string { return ag.id }
// BackendName returns the name of the active backend.
func (ag *Agent) BackendName() string {
if ag.runtime == nil || ag.runtime.Backend == nil {
return ""
}
return ag.runtime.Backend.Name()
}
// ModelName returns the name of the active model.
func (ag *Agent) ModelName() string {
if ag.runtime == nil || ag.runtime.Backend == nil {
return ""
}
return ag.runtime.Backend.Model()
}
// State returns the agent's current execution state.
func (ag *Agent) State() string {
ag.stateMu.RLock()
s := ag.state
ag.stateMu.RUnlock()
return s
}
// SetState sets the agent's execution state and notifies waiters.
func (ag *Agent) SetState(state string) {
ag.stateMu.Lock()
ag.state = state
ag.stateMu.Unlock()
ag.notifyChange()
}
// Reply returns the agent's last assistant response.
func (ag *Agent) Reply() string {
ag.stateMu.RLock()
r := ag.reply
ag.stateMu.RUnlock()
return r
}
// SetReply sets the agent's last response.
func (ag *Agent) SetReply(reply string) {
ag.stateMu.Lock()
ag.reply = reply
ag.stateMu.Unlock()
}
// notifyChange wakes all goroutines waiting on state changes.
func (ag *Agent) notifyChange() {
ag.changeMu.Lock()
ag.changeCond.Broadcast()
ag.changeMu.Unlock()
}
// WaitChange blocks until the agent's state differs from current.
// Returns the new value and true, or ("", false) if ctx is cancelled.
func (ag *Agent) WaitChange(ctx context.Context, field, current string) (string, bool) {
done := make(chan struct{})
context.AfterFunc(ctx, func() {
ag.changeMu.Lock()
ag.changeCond.Broadcast()
ag.changeMu.Unlock()
close(done)
})
ag.changeMu.Lock()
for {
var val string
switch field {
case WatchState:
val = ag.State()
default:
ag.changeMu.Unlock()
return "", false
}
if val != current {
ag.changeMu.Unlock()
return val, true
}
if ctx.Err() != nil {
ag.changeMu.Unlock()
return "", false
}
ag.changeCond.Wait()
if ctx.Err() != nil {
ag.changeMu.Unlock()
return "", false
}
}
}
// InitCond initializes the changeCond. Must be called once after construction.
func (ag *Agent) InitCond() {
ag.changeCond = sync.NewCond(&ag.changeMu)
}
// emit publishes an event on the agent's bus.
func (ag *Agent) emit(ev Event) {
ag.bus.Publish("event", ev)
}
// CWD returns the agent's working directory.
func (ag *Agent) Cwd() string {
ag.stateMu.RLock()
c := ag.cwd
ag.stateMu.RUnlock()
return c
}
// SetCWD sets the agent's working directory (no validation — caller must validate).
func (ag *Agent) SetCwd(dir string) {
ag.stateMu.Lock()
ag.cwd = dir
ag.stateMu.Unlock()
}
// effectiveCwd returns the agent's cwd, falling back to os.Getwd() if empty.
func (ag *Agent) effectiveCwd() string {
if c := ag.Cwd(); c != "" {
return c
}
wd, _ := os.Getwd()
return wd
}
// IsRunning returns true if the agent has an active turn in progress.
func (ag *Agent) IsRunning() bool {
return ag.currentAction.Load() != nil
}
// Interrupt cancels the current in-progress agent turn.
// Returns true if an action was running and was cancelled.
func (ag *Agent) Interrupt(cause error) bool {
if h := ag.currentAction.Load(); h != nil {
h.cancel(cause)
return true
}
return false
}

View File

@ -6,10 +6,9 @@ import (
"fmt" "fmt"
"os" "os"
"path/filepath" "path/filepath"
"slices"
"strconv"
"strings" "strings"
"ollie/agent"
) )
func (a *Session) handleCommand(ctx context.Context, input string) bool { func (a *Session) handleCommand(ctx context.Context, input string) bool {
@ -24,292 +23,21 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
cmd := parts[0] cmd := parts[0]
args := parts[1:] args := parts[1:]
// Try agent-level commands first (they access agent internals directly).
if a.r.HandleCommand(ctx, input) {
return true
}
listFromHandler := func(name string) { listFromHandler := func(name string) {
if h := a.listHandlers[name]; h != nil { if h := a.listHandlers[name]; h != nil {
for _, item := range h() { for _, item := range h() {
a.emit(infoEvent(" " + item)) a.emit(agent.Event{Role: "info", Content: " " + item + "\n"})
} }
} }
} }
type cmdFn func([]string) type cmdFn func([]string)
cmds := map[string]cmdFn{ cmds := map[string]cmdFn{
"/i": func(args []string) {
prompt := strings.Join(args, " ")
if prompt == "" {
a.emit(infoEvent("error: /i requires a prompt"))
return
}
if a.IsRunning() {
a.Inject(prompt)
} else {
go a.Submit(context.Background(), prompt)
}
},
"/irw": func(args []string) {
prompt := strings.Join(args, " ")
if prompt == "" {
a.emit(infoEvent("error: /irw requires a prompt"))
return
}
a.injectRewrite(prompt)
},
"/backend": func(args []string) {
if len(args) == 0 {
a.emit(infoEvent(a.r.runtime.Backend.Name()))
return
}
if a.IsRunning() {
a.emit(infoEvent("error: cannot switch backend while agent is running"))
return
}
be, err := a.r.newBackend(args[0])
if err != nil {
a.emit(infoEvent(fmt.Sprintf("error: failed to switch backend: %v", err)))
return
}
a.r.runtime.Backend = be
a.emit(infoEvent(fmt.Sprintf("switched backend to: %s (model: %s)", be.Name(), be.Model())))
},
"/models": func(args []string) {
models := a.r.runtime.Backend.Models(ctx)
if len(models) == 0 {
a.emit(infoEvent("no models available"))
return
}
slices.Sort(models)
current := a.r.runtime.Backend.Model()
for _, m := range models {
marker := " "
if m == current {
marker = "* "
}
a.emit(infoEvent(marker + m))
}
},
"/model": func(args []string) {
if len(args) == 0 {
a.emit(infoEvent(a.r.runtime.Backend.Model()))
return
}
a.r.runtime.Backend.SetModel(args[0])
a.emit(infoEvent("switched model to: " + args[0]))
},
"/maxsteps": func(args []string) {
if len(args) == 0 {
if a.r.runtime.MaxSteps == 0 {
a.emit(infoEvent("maxsteps: unlimited"))
} else {
a.emit(infoEvent(fmt.Sprintf("maxsteps: %d", a.r.runtime.MaxSteps)))
}
return
}
n, err := strconv.Atoi(args[0])
if err != nil || n < 0 {
a.emit(infoEvent("error: maxsteps requires a non-negative integer (0 = unlimited)"))
return
}
a.r.runtime.MaxSteps = n
if n == 0 {
a.emit(infoEvent("maxsteps: unlimited"))
} else {
a.emit(infoEvent(fmt.Sprintf("maxsteps: %d", n)))
}
},
"/agents": func(args []string) {
seen := make(map[string]bool)
found := false
for _, dir := range AgentsDirs() {
entries, err := os.ReadDir(dir)
if err != nil {
continue
}
for _, e := range entries {
if e.IsDir() || !strings.HasSuffix(e.Name(), ".json") {
continue
}
name := strings.TrimSuffix(e.Name(), ".json")
if seen[name] {
continue
}
seen[name] = true
marker := " "
if name == a.r.agentName {
marker = "* "
}
a.emit(infoEvent(marker + name))
found = true
}
}
if !found {
a.emit(infoEvent("no agents found"))
}
},
"/agent": func(args []string) {
if len(args) == 0 {
a.emit(infoEvent("active agent: " + a.r.agentName))
return
}
if a.IsRunning() {
a.emit(infoEvent("error: cannot switch agent while agent is running"))
return
}
name := args[0]
cfgPath := AgentConfigPath(a.r.agentsDir, name)
f, err := os.Open(cfgPath)
if err != nil {
a.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err)))
return
}
cfg, err := Load(f)
f.Close()
if err != nil {
a.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err)))
return
}
d := a.r.newDispatcher()
env := []string{"OLLIE_SESSION_ID=" + a.id, "OLLIE_UNAME=" + a.r.id}
env = append(env, a.r.promptEnvExtra...)
rt := BuildRuntime(cfg, d, a.r.cwd, env, a.r.baseLayers...)
if rt.CfgBackend != "" {
newBe, err := a.r.newBackend(rt.CfgBackend)
if err != nil {
a.emit(infoEvent(fmt.Sprintf("error: backend %q: %v", rt.CfgBackend, err)))
return
}
if rt.CfgModel != "" {
newBe.SetModel(rt.CfgModel)
}
rt.Backend = newBe
} else {
rt.Backend = a.r.runtime.Backend
if rt.CfgModel != "" {
rt.Backend.SetModel(rt.CfgModel)
}
}
a.r.runtime = rt
a.r.agentName = name
a.r.history = nil
a.pushSessionEnv()
a.r.notifyChange()
for _, msg := range rt.Messages {
a.emit(infoEvent(msg))
}
a.emit(infoEvent("agent: " + name))
},
"/compact": func(args []string) {
if a.IsRunning() {
a.emit(infoEvent("error: cannot compact while agent is running"))
return
}
if a.r.history == nil {
a.emit(infoEvent("nothing to compact"))
return
}
snapshot := a.r.history.PreCompactionSnapshot()
a.r.SetState("compacting")
n, err := a.r.runCompact(ctx, "manual")
a.r.SetState("idle")
if err != nil {
a.emit(infoEvent("compact error: " + err.Error()))
return
}
if n == 0 {
a.emit(infoEvent("nothing to compact"))
return
}
if a.sessionsDir != "" && a.id != "" {
histPath := a.activeSessionPath(a.id, ".compaction.jsonl")
if err := os.MkdirAll(filepath.Dir(histPath), 0700); err != nil {
a.emit(infoEvent("compaction history save: " + err.Error()))
} else if data, err := json.Marshal(snapshot); err == nil {
f, err := os.OpenFile(histPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0600)
if err == nil {
f.Write(append(data, '\n')) //nolint:errcheck
f.Close() //nolint:errcheck
}
}
}
a.emit(infoEvent(fmt.Sprintf("compacted %d messages", n)))
a.saveSession()
},
"/context": func(args []string) {
if a.r.history == nil {
a.emit(infoEvent("no active session"))
return
}
ctxLen := a.r.runtime.Backend.ContextLength(ctx)
if ctxLen <= 0 {
ctxLen = defaultContextLength
}
estimated := a.r.history.estimateTokens()
pct := estimated * 100 / ctxLen
a.emit(infoEvent(fmt.Sprintf("~%d / %d tokens (%d%%)", estimated, ctxLen, pct)))
a.emit(infoEvent(strings.TrimRight(a.r.history.contextDebug(), "\n")))
},
"/cost": func(args []string) {
if a.r.history == nil {
a.emit(infoEvent("no active session"))
return
}
a.emit(infoEvent(fmt.Sprintf("last=$%.4f session=$%.4f",
a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD)))
},
"/usage": func(args []string) {
if a.r.history == nil {
a.emit(infoEvent("no active session"))
return
}
ctxLen := a.r.runtime.Backend.ContextLength(ctx)
if ctxLen <= 0 {
ctxLen = defaultContextLength
}
estimated := a.r.history.estimateTokens()
pct := estimated * 100 / ctxLen
usageStr := fmt.Sprintf("~%d / %d tokens (%d%%) | %d in, %d out, %d requests",
estimated, ctxLen, pct,
a.r.history.TotalInputTokens, a.r.history.TotalOutputTokens,
a.r.history.TotalRequests)
if a.r.history.Estimated {
usageStr += " [estimated]"
}
a.emit(infoEvent(usageStr))
},
"/history": func(args []string) {
if a.r.history == nil {
a.emit(infoEvent("no active session"))
return
}
for _, msg := range a.r.history.history() {
preview := msg.Content
if len(preview) > 200 {
preview = preview[:200] + "..."
}
a.emit(infoEvent(fmt.Sprintf("[%s] %s", msg.Role, preview)))
}
},
"/clear": func(args []string) {
if a.IsRunning() {
a.emit(infoEvent("error: cannot clear while agent is running"))
return
}
a.r.history = nil
a.emit(infoEvent("cleared"))
},
"/sessions": func(args []string) { "/sessions": func(args []string) {
allFlag := len(args) > 0 && args[0] == "-a" allFlag := len(args) > 0 && args[0] == "-a"
type sessionFile struct { type sessionFile struct {
@ -343,7 +71,7 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
if readErr != nil { if readErr != nil {
continue continue
} }
var ps PersistedAgent var ps agent.PersistedAgent
if json.Unmarshal(data, &ps) != nil { if json.Unmarshal(data, &ps) != nil {
continue continue
} }
@ -364,77 +92,73 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
if len(goal) > 60 { if len(goal) > 60 {
goal = goal[:60] + "..." goal = goal[:60] + "..."
} }
a.emit(infoEvent(marker + fmt.Sprintf("%-24s [%s] %q", file.id, ps.Agent, goal))) a.emit(agent.Event{Role: "info", Content: marker + fmt.Sprintf("%-24s [%s] %q", file.id, ps.Agent, goal) + "\n"})
found = true found = true
} }
if !found { if !found {
a.emit(infoEvent("no sessions for " + cwd)) a.emit(agent.Event{Role: "info", Content: "no sessions for " + cwd + "\n"})
} }
}, },
"/save": func(args []string) { "/save": func(args []string) {
if a.r.history == nil { if !a.r.HasHistory() {
a.emit(infoEvent("error: no active session")) a.emit(agent.Event{Role: "info", Content: "error: no active session\n"})
return return
} }
if len(args) == 0 { if len(args) == 0 {
a.emit(infoEvent("error: /save requires a name")) a.emit(agent.Event{Role: "info", Content: "error: /save requires a name\n"})
return return
} }
name := args[0] name := args[0]
path := a.sessionsDir + "/" + name + ".json" path := a.sessionsDir + "/" + name + ".json"
if err := a.r.history.saveTo(path, name, a.r.agentName, a.CWD()); err != nil { if err := a.r.SaveTo(path, name, a.CWD()); err != nil {
a.emit(infoEvent("error: " + err.Error())) a.emit(agent.Event{Role: "info", Content: "error: " + err.Error() + "\n"})
return return
} }
a.emit(infoEvent("saved: " + path)) a.emit(agent.Event{Role: "info", Content: "saved: " + path + "\n"})
}, },
"/resume": func(args []string) { "/resume": func(args []string) {
if len(args) == 0 { if len(args) == 0 {
a.emit(infoEvent("error: /resume requires a session id or name")) a.emit(agent.Event{Role: "info", Content: "error: /resume requires a session id or name\n"})
return return
} }
if a.IsRunning() { if a.IsRunning() {
a.emit(infoEvent("error: cannot resume while agent is running")) a.emit(agent.Event{Role: "info", Content: "error: cannot resume while agent is running\n"})
return return
} }
name := args[0] name := args[0]
path := a.sessionsDir + "/" + name + ".json" path := a.sessionsDir + "/" + name + ".json"
data, err := os.ReadFile(path) data, err := os.ReadFile(path)
if err != nil { if err != nil {
a.emit(infoEvent(fmt.Sprintf("error: %v", err))) a.emit(agent.Event{Role: "info", Content: fmt.Sprintf("error: %v\n", err)})
return return
} }
var ps PersistedAgent var ps agent.PersistedAgent
if err := json.Unmarshal(data, &ps); err != nil { if err := json.Unmarshal(data, &ps); err != nil {
a.emit(infoEvent(fmt.Sprintf("error: %v", err))) a.emit(agent.Event{Role: "info", Content: fmt.Sprintf("error: %v\n", err)})
return return
} }
a.r.history = RestoreHistory(&ps) a.r.Restore(&ps)
a.emit(infoEvent(fmt.Sprintf("resumed session %s (%d messages)", name, len(ps.Messages)))) a.emit(agent.Event{Role: "info", Content: fmt.Sprintf("resumed session %s (%d messages)\n", name, len(ps.Messages))})
}, },
"/cwd": func(args []string) { "/cwd": func(args []string) {
if len(args) == 0 { if len(args) == 0 {
a.emit(infoEvent("cwd: " + a.CWD())) a.emit(agent.Event{Role: "info", Content: "cwd: " + a.CWD() + "\n"})
return return
} }
dir := strings.Join(args, " ") dir := strings.Join(args, " ")
if err := a.SetCWD(dir); err != nil { if err := a.SetCWD(dir); err != nil {
a.emit(infoEvent("error: " + err.Error())) a.emit(agent.Event{Role: "info", Content: "error: " + err.Error() + "\n"})
return return
} }
a.emit(infoEvent("cwd: " + dir)) a.emit(agent.Event{Role: "info", Content: "cwd: " + dir + "\n"})
}, },
"/skills": func(args []string) { listFromHandler("skills") }, "/skills": func(args []string) { listFromHandler("skills") },
"/tools": func(args []string) { listFromHandler("tools") }, "/tools": func(args []string) { listFromHandler("tools") },
"/sp": func(args []string) {
a.emit(infoEvent(a.r.runtime.Preamble))
},
"/help": func(args []string) { "/help": func(args []string) {
lines := []string{ lines := []string{
"Available commands:", "Available commands:",
@ -452,21 +176,17 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
" /cwd [path] - show or change working directory", " /cwd [path] - show or change working directory",
" /i <prompt> - inject prompt into the running turn", " /i <prompt> - inject prompt into the running turn",
" /irw <prompt> - rewrite the pending inject", " /irw <prompt> - rewrite the pending inject",
" /queued [pop|clear] - manage queued prompts",
" /compact - summarize conversation and compact context", " /compact - summarize conversation and compact context",
" /context - show context size and message breakdown", " /context - show context size and message breakdown",
" /cost - show last turn and session cost", " /cost - show last turn and session cost",
" /usage - show token usage and context percentage", " /usage - show token usage and context percentage",
" /history - dump bounded message history", " /history - dump bounded message history",
" /clear - clear session", " /clear - clear session",
" /kill - kill session",
" /rn <name> - rename session",
" /sp - show rendered system prompt", " /sp - show rendered system prompt",
" /help - show this help", " /help - show this help",
" !<cmd> - run shell command",
} }
for _, l := range lines { for _, l := range lines {
a.emit(infoEvent(l)) a.emit(agent.Event{Role: "info", Content: l + "\n"})
} }
}, },
} }
@ -475,7 +195,7 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool {
if !ok { if !ok {
return false return false
} }
a.emit(infoEvent("")) a.emit(agent.Event{Role: "info", Content: "\n"})
fn(args) fn(args)
return true return true
} }

View File

@ -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)
}
}
}
}

View File

@ -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")
}
}

File diff suppressed because it is too large Load Diff

View File

@ -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)
}
}

View File

@ -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 }

View File

@ -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")
}
}

View File

@ -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) {}

View File

@ -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 }

View File

@ -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
}
}
}

View File

@ -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)])
}
}
}

View File

@ -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)
}
}
}

File diff suppressed because it is too large Load Diff

View File

@ -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)
}
}
}