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