diff --git a/go.mod b/go.mod index c0d712d..0737380 100644 --- a/go.mod +++ b/go.mod @@ -5,3 +5,5 @@ go 1.25.6 require gopkg.in/yaml.v3 v3.0.1 require github.com/aymanbagabas/go-udiff v0.4.1 + +require github.com/simonfxr/pubsub v0.0.5 diff --git a/go.sum b/go.sum index 52afe9d..e289754 100644 --- a/go.sum +++ b/go.sum @@ -1,6 +1,13 @@ github.com/aymanbagabas/go-udiff v0.4.1 h1:OEIrQ8maEeDBXQDoGCbbTTXYJMYRCRO1fnodZ12Gv5o= github.com/aymanbagabas/go-udiff v0.4.1/go.mod h1:0L9PGwj20lrtmEMeyw4WKJ/TMyDtvAoK9bf2u/mNo3w= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/simonfxr/pubsub v0.0.5 h1:DJfvFoglqGvwJriIOC5NI5um34n2YX8KAU5+7jv768w= +github.com/simonfxr/pubsub v0.0.5/go.mod h1:bQ+B2NEEHZ08VY/0xVFGaGL1KMJ7cl3yLaXe8xvyC24= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/pkg/agent/commands.go b/pkg/agent/commands.go index 02aec8f..30f4113 100644 --- a/pkg/agent/commands.go +++ b/pkg/agent/commands.go @@ -13,9 +13,9 @@ import ( "ollie/pkg/config" ) -func (s *agent) handleCommand(ctx context.Context, input string, handler EventHandler) bool { +func (s *agent) handleCommand(ctx context.Context, input string) bool { if strings.HasPrefix(input, "!") { - handler(infoEvent("")) + s.emit(infoEvent("")) cmdStr := strings.TrimSpace(input[1:]) if cmdStr == "" { return true @@ -30,10 +30,10 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa s.envMu.RUnlock() o, err := shellCmd.CombinedOutput() if err != nil { - handler(infoEvent("error: " + err.Error())) + s.emit(infoEvent("error: " + err.Error())) } if len(o) > 0 { - handler(infoEvent(strings.TrimRight(string(o), "\n"))) + s.emit(infoEvent(strings.TrimRight(string(o), "\n"))) } return true } @@ -57,12 +57,12 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa } entries, err := os.ReadDir(mount + "/" + subdir) if err != nil { - handler(infoEvent(fmt.Sprintf("%s: %v", cmd, err))) + s.emit(infoEvent(fmt.Sprintf("%s: %v", cmd, err))) return } for _, e := range entries { if !e.IsDir() { - handler(infoEvent(" " + e.Name())) + s.emit(infoEvent(" " + e.Name())) } } } @@ -72,16 +72,20 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa "/i": func(args []string) { prompt := strings.Join(args, " ") if prompt == "" { - handler(infoEvent("error: /i requires a prompt")) + s.emit(infoEvent("error: /i requires a prompt")) return } - s.Inject(prompt) + if s.IsRunning() { + s.Inject(prompt) + } else { + go s.Submit(context.Background(), prompt) + } }, "/irw": func(args []string) { prompt := strings.Join(args, " ") if prompt == "" { - handler(infoEvent("error: /irw requires a prompt")) + s.emit(infoEvent("error: /irw requires a prompt")) return } s.injectRewrite(prompt) @@ -89,26 +93,26 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa "/backend": func(args []string) { if len(args) == 0 { - handler(infoEvent(s.cfg.Backend.Name())) + s.emit(infoEvent(s.cfg.Backend.Name())) return } if s.IsRunning() { - handler(infoEvent("error: cannot switch backend while agent is running")) + s.emit(infoEvent("error: cannot switch backend while agent is running")) return } be, err := s.newBackend(args[0]) if err != nil { - handler(infoEvent(fmt.Sprintf("error: failed to switch backend: %v", err))) + s.emit(infoEvent(fmt.Sprintf("error: failed to switch backend: %v", err))) return } s.cfg.Backend = be - handler(infoEvent(fmt.Sprintf("switched backend to: %s (model: %s)", be.Name(), be.Model()))) + s.emit(infoEvent(fmt.Sprintf("switched backend to: %s (model: %s)", be.Name(), be.Model()))) }, "/models": func(args []string) { models := s.cfg.Backend.Models(ctx) if len(models) == 0 { - handler(infoEvent("no models available")) + s.emit(infoEvent("no models available")) return } slices.Sort(models) @@ -118,38 +122,38 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa if m == current { marker = "* " } - handler(infoEvent(marker + m)) + s.emit(infoEvent(marker + m)) } }, "/model": func(args []string) { if len(args) == 0 { - handler(infoEvent(s.cfg.Backend.Model())) + s.emit(infoEvent(s.cfg.Backend.Model())) return } s.cfg.Backend.SetModel(args[0]) - handler(infoEvent("switched model to: " + args[0])) + s.emit(infoEvent("switched model to: " + args[0])) }, "/maxsteps": func(args []string) { if len(args) == 0 { if s.cfg.MaxSteps == 0 { - handler(infoEvent("maxsteps: unlimited")) + s.emit(infoEvent("maxsteps: unlimited")) } else { - handler(infoEvent(fmt.Sprintf("maxsteps: %d", s.cfg.MaxSteps))) + s.emit(infoEvent(fmt.Sprintf("maxsteps: %d", s.cfg.MaxSteps))) } return } n, err := strconv.Atoi(args[0]) if err != nil || n < 0 { - handler(infoEvent("error: maxsteps requires a non-negative integer (0 = unlimited)")) + s.emit(infoEvent("error: maxsteps requires a non-negative integer (0 = unlimited)")) return } s.cfg.MaxSteps = n if n == 0 { - handler(infoEvent("maxsteps: unlimited")) + s.emit(infoEvent("maxsteps: unlimited")) } else { - handler(infoEvent(fmt.Sprintf("maxsteps: %d", n))) + s.emit(infoEvent(fmt.Sprintf("maxsteps: %d", n))) } }, @@ -174,35 +178,35 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa if name == s.agentName { marker = "* " } - handler(infoEvent(marker + name)) + s.emit(infoEvent(marker + name)) found = true } } if !found { - handler(infoEvent("no agents found")) + s.emit(infoEvent("no agents found")) } }, "/agent": func(args []string) { if len(args) == 0 { - handler(infoEvent("active agent: " + s.agentName)) + s.emit(infoEvent("active agent: " + s.agentName)) return } if s.IsRunning() { - handler(infoEvent("error: cannot switch agent while agent is running")) + s.emit(infoEvent("error: cannot switch agent while agent is running")) return } name := args[0] cfgPath := AgentConfigPath(s.agentsDir, name) f, err := os.Open(cfgPath) if err != nil { - handler(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err))) + s.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err))) return } cfg, err := config.Load(f) f.Close() if err != nil { - handler(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err))) + s.emit(infoEvent(fmt.Sprintf("error: agent %q: %v", name, err))) return } d := s.newDispatcher() @@ -210,7 +214,7 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa if env.Backend != "" { newBe, err := s.newBackend(env.Backend) if err != nil { - handler(infoEvent(fmt.Sprintf("error: backend %q: %v", env.Backend, err))) + s.emit(infoEvent(fmt.Sprintf("error: backend %q: %v", env.Backend, err))) return } if env.Model != "" { @@ -231,30 +235,30 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa s.session = nil s.pushSessionEnv() for _, msg := range env.Messages { - handler(infoEvent(msg)) + s.emit(infoEvent(msg)) } - handler(infoEvent("agent: " + name)) + s.emit(infoEvent("agent: " + name)) }, "/compact": func(args []string) { if s.IsRunning() { - handler(infoEvent("error: cannot compact while agent is running")) + s.emit(infoEvent("error: cannot compact while agent is running")) return } if s.session == nil { - handler(infoEvent("nothing to compact")) + s.emit(infoEvent("nothing to compact")) return } snapshot := s.session.PreCompactionSnapshot() s.setState("compacting") - n, err := s.runCompact(ctx, "manual", handler) + n, err := s.runCompact(ctx, "manual") s.setState("idle") if err != nil { - handler(infoEvent("compact error: " + err.Error())) + s.emit(infoEvent("compact error: " + err.Error())) return } if n == 0 { - handler(infoEvent("nothing to compact")) + s.emit(infoEvent("nothing to compact")) return } if s.sessionsDir != "" && s.sessionID != "" { @@ -267,13 +271,13 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa } } } - handler(infoEvent(fmt.Sprintf("compacted %d messages", n))) + s.emit(infoEvent(fmt.Sprintf("compacted %d messages", n))) s.saveSession() }, "/context": func(args []string) { if s.session == nil { - handler(infoEvent("no active session")) + s.emit(infoEvent("no active session")) return } ctxLen := s.cfg.Backend.ContextLength(ctx) @@ -282,22 +286,22 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa } estimated := s.session.estimateTokens() pct := estimated * 100 / ctxLen - handler(infoEvent(fmt.Sprintf("~%d / %d tokens (%d%%)", estimated, ctxLen, pct))) - handler(infoEvent(strings.TrimRight(s.session.contextDebug(), "\n"))) + s.emit(infoEvent(fmt.Sprintf("~%d / %d tokens (%d%%)", estimated, ctxLen, pct))) + s.emit(infoEvent(strings.TrimRight(s.session.contextDebug(), "\n"))) }, "/cost": func(args []string) { if s.session == nil { - handler(infoEvent("no active session")) + s.emit(infoEvent("no active session")) return } - handler(infoEvent(fmt.Sprintf("last=$%.4f session=$%.4f", + s.emit(infoEvent(fmt.Sprintf("last=$%.4f session=$%.4f", s.session.LastTurnCostUSD, s.session.SessionCostUSD))) }, "/usage": func(args []string) { if s.session == nil { - handler(infoEvent("no active session")) + s.emit(infoEvent("no active session")) return } ctxLen := s.cfg.Backend.ContextLength(ctx) @@ -313,12 +317,12 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa if s.session.Estimated { usageStr += " [estimated]" } - handler(infoEvent(usageStr)) + s.emit(infoEvent(usageStr)) }, "/history": func(args []string) { if s.session == nil { - handler(infoEvent("no active session")) + s.emit(infoEvent("no active session")) return } for _, msg := range s.session.history() { @@ -326,23 +330,23 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa if len(preview) > 200 { preview = preview[:200] + "..." } - handler(infoEvent(fmt.Sprintf("[%s] %s", msg.Role, preview))) + s.emit(infoEvent(fmt.Sprintf("[%s] %s", msg.Role, preview))) } }, "/clear": func(args []string) { if s.IsRunning() { - handler(infoEvent("error: cannot clear while agent is running")) + s.emit(infoEvent("error: cannot clear while agent is running")) return } s.session = nil - handler(infoEvent("cleared")) + s.emit(infoEvent("cleared")) }, "/sessions": func(args []string) { entries, err := os.ReadDir(s.sessionsDir) if err != nil { - handler(infoEvent(fmt.Sprintf("sessions: %v", err))) + s.emit(infoEvent(fmt.Sprintf("sessions: %v", err))) return } found := false @@ -373,32 +377,32 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa label = fmt.Sprintf("%-24s [%s] %q", id, ps.Agent, goal) } } - handler(infoEvent(marker + label)) + s.emit(infoEvent(marker + label)) found = true } if !found { - handler(infoEvent("no sessions found in " + s.sessionsDir)) + s.emit(infoEvent("no sessions found in " + s.sessionsDir)) } }, "/cwd": func(args []string) { if len(args) == 0 { - handler(infoEvent("cwd: " + s.CWD())) + s.emit(infoEvent("cwd: " + s.CWD())) return } dir := strings.Join(args, " ") if err := s.SetCWD(dir); err != nil { - handler(infoEvent("error: " + err.Error())) + s.emit(infoEvent("error: " + err.Error())) return } - handler(infoEvent("cwd: " + dir)) + s.emit(infoEvent("cwd: " + dir)) }, "/skills": func(args []string) { listMountDir("sk") }, "/tools": func(args []string) { listMountDir("t") }, "/sp": func(args []string) { - handler(infoEvent(s.cfg.preamble)) + s.emit(infoEvent(s.cfg.preamble)) }, "/help": func(args []string) { @@ -430,7 +434,7 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa " ! - run shell command", } for _, l := range lines { - handler(infoEvent(l)) + s.emit(infoEvent(l)) } }, } @@ -439,7 +443,7 @@ func (s *agent) handleCommand(ctx context.Context, input string, handler EventHa if !ok { return false } - handler(infoEvent("")) + s.emit(infoEvent("")) fn(args) return true } diff --git a/pkg/agent/core.go b/pkg/agent/core.go index 7d2b380..3f0e4b6 100644 --- a/pkg/agent/core.go +++ b/pkg/agent/core.go @@ -17,6 +17,7 @@ import ( "sync/atomic" "time" + "github.com/simonfxr/pubsub" "ollie/pkg/backend" "ollie/pkg/config" olog "ollie/pkg/log" @@ -267,6 +268,7 @@ type agent struct { currentAction atomic.Pointer[actionHandle] toolCallCount atomic.Int64 fifo PromptFIFO + bus *pubsub.Bus pendingInject atomic.Pointer[string] mu sync.RWMutex state string // "idle", "thinking", "calling: " @@ -428,6 +430,7 @@ func NewAgentCore(cfg AgentCoreConfig) Core { newDispatcher: cfg.NewDispatcher, newBackend: cfg.NewBackend, state: "idle", + bus: pubsub.NewBus(), } a.changeCond = sync.NewCond(&a.changeMu) a.pushSessionEnv() @@ -637,7 +640,7 @@ func (s *agent) autoWarnLimit(ctx context.Context) int { // spawnContext assembles the agent context injected at each session refresh // point (session start, post-clear, post-compaction). It combines the // agent-specific prompt with any agentSpawn hook output. -func (s *agent) spawnContext(ctx context.Context, handler EventHandler) string { +func (s *agent) spawnContext(ctx context.Context) string { result := s.hooks.Run(ctx, HookAgentSpawn, map[string]string{ "session_id": s.sessionID, "agent": s.agentName, @@ -645,7 +648,7 @@ func (s *agent) spawnContext(ctx context.Context, handler EventHandler) string { "model": s.cfg.Backend.Model(), }, s.log) if result.Warning != "" { - handler(infoEvent(result.Warning)) + s.emit(infoEvent(result.Warning)) } var parts []string if result.Context != "" { @@ -657,14 +660,14 @@ func (s *agent) spawnContext(ctx context.Context, handler EventHandler) string { // runCompact executes a full compaction cycle: pre-hook, compact, spawn-context // re-injection, post-hook. Returns (n compacted, error). Returns (0, nil) if // the pre-hook blocked or there was nothing to compact. Caller manages setState. -func (s *agent) runCompact(ctx context.Context, trigger string, handler EventHandler) (int, error) { +func (s *agent) runCompact(ctx context.Context, trigger string) (int, error) { payload := map[string]string{"session_id": s.sessionID, "trigger": trigger, "cwd": s.CWD()} pre := s.hooks.Run(ctx, HookPreCompact, payload, s.log) if pre.Warning != "" { - handler(infoEvent(pre.Warning)) + s.emit(infoEvent(pre.Warning)) } if pre.Blocked { - handler(infoEvent("compact cancelled by hook")) + s.emit(infoEvent("compact cancelled by hook")) return 0, nil } if pre.Context != "" { @@ -677,13 +680,13 @@ func (s *agent) runCompact(ctx context.Context, trigger string, handler EventHan if n > 0 { s.auditLog.Debug("compact: removed %d messages trigger=%s session=%s", n, trigger, s.sessionID) s.warnedContext = false - if sc := s.spawnContext(ctx, handler); sc != "" { + if sc := s.spawnContext(ctx); sc != "" { s.session.appendUserMessage(sc) } } post := s.hooks.Run(ctx, HookPostCompact, payload, s.log) if post.Warning != "" { - handler(infoEvent(post.Warning)) + s.emit(infoEvent(post.Warning)) } if post.Context != "" { s.session.appendUserMessage(post.Context) @@ -739,6 +742,24 @@ func (s *agent) injectRewrite(prompt string) { func (s *agent) Queue(prompt string) { s.log.Debug("Queue() len=%d", len(prompt)) s.fifo.Push(prompt) + s.bus.Publish("queued", prompt) + if !s.IsRunning() { + go s.drainQueue() + } +} + +func (s *agent) drainQueue() { + if prompt, ok := s.fifo.Pop(); ok { + s.Submit(context.Background(), prompt) + } +} + +func (s *agent) Bus() *pubsub.Bus { + return s.bus +} + +func (s *agent) emit(ev Event) { + s.bus.Publish("event", ev) } func (s *agent) PopQueue() (string, bool) { @@ -852,13 +873,13 @@ func firstSentence(s string) string { } // Submit implements Core. It processes one line of user input: slash commands -// and shell shortcuts are dispatched immediately via handler; any other input -// starts an agent turn that streams events to handler. If a turn is already +// and shell shortcuts are dispatched immediately; any other input +// starts an agent turn that streams events to the bus. If a turn is already // in progress the prompt is queued as an in-stream interruption instead. // // Continuations (post-turn hook context, unconsumed inject, FIFO drain) are // handled via an explicit loop rather than recursion to avoid stack growth. -func (s *agent) Submit(ctx context.Context, input string, handler EventHandler) { +func (s *agent) Submit(ctx context.Context, input string) { defer func() { if r := recover(); r != nil { s.log.Error("panic: %v\n%s", r, debug.Stack()) @@ -866,7 +887,7 @@ func (s *agent) Submit(ctx context.Context, input string, handler EventHandler) a.cancel(fmt.Errorf("%v", r)) } s.setState("idle") - handler(Event{Role: "error", Content: fmt.Sprintf("%v", r)}) + s.emit(Event{Role: "error", Content: fmt.Sprintf("%v", r)}) } }() s.log.Debug("Submit() input_len=%d running=%v", len(input), s.IsRunning()) @@ -878,7 +899,7 @@ func (s *agent) Submit(ctx context.Context, input string, handler EventHandler) // the submit lock. Handle them before acquiring submitMu so they // don't block behind a long-running turn or command. if s.IsRunning() { - if s.handleCommand(ctx, input, handler) { + if s.handleCommand(ctx, input) { return } s.fifo.Push(input) @@ -890,7 +911,7 @@ func (s *agent) Submit(ctx context.Context, input string, handler EventHandler) s.submitMu.Lock() defer s.submitMu.Unlock() - if s.handleCommand(ctx, input, handler) { + if s.handleCommand(ctx, input) { return } if s.IsRunning() { @@ -899,14 +920,14 @@ func (s *agent) Submit(ctx context.Context, input string, handler EventHandler) } for input != "" && ctx.Err() == nil { - input = s.executeTurn(ctx, input, handler) + input = s.executeTurn(ctx, input) } } // executeTurn runs a single agent turn and returns the next prompt to execute, // or "" if there is nothing more to do. -func (s *agent) executeTurn(ctx context.Context, input string, handler EventHandler) string { - handler(Event{Role: "user", Content: input}) +func (s *agent) executeTurn(ctx context.Context, input string) string { + s.emit(Event{Role: "user", Content: input}) hookResult := s.hooks.Run(ctx, HookPreTurn, map[string]string{ "session_id": s.sessionID, @@ -914,11 +935,11 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand "prompt": input, }, s.log) if hookResult.Blocked { - handler(infoEvent("hook blocked prompt")) + s.emit(infoEvent("hook blocked prompt")) return "" } if hookResult.Warning != "" { - handler(infoEvent(hookResult.Warning)) + s.emit(infoEvent(hookResult.Warning)) } if hookResult.Context != "" { input += "\n" + hookResult.Context @@ -935,11 +956,11 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand if s.session == nil { for _, msg := range s.startupMessages { s.log.Debug("startup: %s", msg) - handler(infoEvent(msg)) + s.emit(infoEvent(msg)) } s.startupMessages = nil s.session = newSession(input) - if sc := s.spawnContext(ctx, handler); sc != "" { + if sc := s.spawnContext(ctx); sc != "" { s.session.appendUserMessage(sc) } s.session.appendUserMessage(input) @@ -984,12 +1005,12 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand s.notifyChange() if maxCostStr := os.Getenv("OLLIE_MAX_SESSION_COST"); maxCostStr != "" { if limit, ferr := strconv.ParseFloat(maxCostStr, 64); ferr == nil && limit > 0 && s.session.SessionCostUSD >= limit { - handler(infoEvent(fmt.Sprintf("spending cap $%.2f reached — stopping", limit))) + s.emit(infoEvent(fmt.Sprintf("spending cap $%.2f reached — stopping", limit))) s.Interrupt(ErrInterrupted) } } } - handler(ev) + s.emit(ev) } s.cfg.PopInject = func() string { if p := s.pendingInject.Swap(nil); p != nil { @@ -1028,7 +1049,7 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand } emit(s.cfg, Event{Role: "info", Content: "auto-compacting context...\n"}) s.setState("compacting") - if _, err := s.runCompact(ctx, "auto", handler); err != nil { + if _, err := s.runCompact(ctx, "auto"); err != nil { panic(fmt.Sprintf("mid-turn auto-compact: %v", err)) } s.setState("thinking") @@ -1040,7 +1061,7 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand if compactLimit := s.autoCompactLimit(ctx); compactLimit > 0 && tokens >= compactLimit { emit(s.cfg, Event{Role: "info", Content: "auto-compacting context...\n"}) s.setState("compacting") - if _, err := s.runCompact(ctx, "auto", handler); err != nil { + if _, err := s.runCompact(ctx, "auto"); err != nil { panic(fmt.Sprintf("auto-compact: %v", err)) } s.setState("thinking") @@ -1059,7 +1080,7 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand if maxCostStr := os.Getenv("OLLIE_MAX_SESSION_COST"); maxCostStr != "" && s.session != nil { if limit, err := strconv.ParseFloat(maxCostStr, 64); err == nil && limit > 0 { if s.session.SessionCostUSD >= limit { - handler(Event{Role: "error", Content: fmt.Sprintf("spending cap $%.2f reached (session total $%.4f)", limit, s.session.SessionCostUSD)}) + s.emit(Event{Role: "error", Content: fmt.Sprintf("spending cap $%.2f reached (session total $%.4f)", limit, s.session.SessionCostUSD)}) s.setState("idle") actCancel(nil) s.currentAction.CompareAndSwap(handle, nil) @@ -1099,7 +1120,7 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand s.session.messages = snapMessages emit(s.cfg, Event{Role: "info", Content: "context overflow — compacting and retrying...\n"}) s.setState("compacting") - if _, cerr := s.runCompact(ctx, "overflow", handler); cerr != nil { + if _, cerr := s.runCompact(ctx, "overflow"); cerr != nil { break } s.session.appendUserMessage(input) @@ -1131,7 +1152,7 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand s.saveSession() return "" } - handler(Event{Role: "error", Content: err.Error()}) + s.emit(Event{Role: "error", Content: err.Error()}) // Drain one FIFO item — the turnError hook may have queued a recovery prompt. if next, ok := s.fifo.Pop(); ok { return next @@ -1144,7 +1165,7 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand "cwd": s.CWD(), }, s.log) if stopResult.Warning != "" { - handler(infoEvent(stopResult.Warning)) + s.emit(infoEvent(stopResult.Warning)) } if !stopResult.Blocked && stopResult.Context != "" && s.session != nil { s.session.appendUserMessage(stopResult.Context) @@ -1153,7 +1174,7 @@ func (s *agent) executeTurn(ctx context.Context, input string, handler EventHand if s.session != nil { s.session.recordTurnCost(s.cfg.Backend.Model()) if s.session.LastTurnCostUSD > 0 { - handler(Event{Role: "info", Content: fmt.Sprintf("costLast=$%.4f\n", s.session.LastTurnCostUSD)}) + s.emit(Event{Role: "info", Content: fmt.Sprintf("costLast=$%.4f\n", s.session.LastTurnCostUSD)}) } s.auditLog.Debug("turn: end reply=%s cost=$%.4f session_total=$%.4f session=%s", auditTruncate(s.reply), s.session.LastTurnCostUSD, s.session.SessionCostUSD, s.sessionID) diff --git a/pkg/agent/core_test.go b/pkg/agent/core_test.go index e9c8b43..4db6e21 100644 --- a/pkg/agent/core_test.go +++ b/pkg/agent/core_test.go @@ -8,6 +8,7 @@ import ( "path/filepath" "strings" "sync" + "sync/atomic" "testing" "time" @@ -114,11 +115,15 @@ func newCore(t *testing.T, be backend.Backend, hooks Hooks) *agent { func collectEvents(ctx context.Context, c Core, input string) []Event { var mu sync.Mutex var evs []Event - c.Submit(ctx, input, func(ev Event) { + sub := c.Bus().Subscribe("event", func(ev Event) { mu.Lock() evs = append(evs, ev) mu.Unlock() }) + c.Submit(ctx, input) + c.Bus().Unsubscribe(sub) + mu.Lock() + defer mu.Unlock() return evs } @@ -190,7 +195,7 @@ func TestSubmit_StateTransitions(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -313,20 +318,26 @@ func TestSubmit_PostTurnHookContinue(t *testing.T) { // --- FIFO drain --- func TestSubmit_FIFODrain(t *testing.T) { - callCount := 0 + var callCount atomic.Int32 be := defaultBE() be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { - callCount++ + callCount.Add(1) return textStream("ok"), nil } c := newCore(t, be, nil) + c.Queue("second") c.Queue("third") + c.Submit(context.Background(), "first") - collectEvents(context.Background(), c, "first") + // Wait for async drain goroutines to complete. + deadline := time.Now().Add(2 * time.Second) + for callCount.Load() < 3 && time.Now().Before(deadline) { + time.Sleep(5 * time.Millisecond) + } - if callCount != 3 { - t.Errorf("backend called %d times; want 3 (first + second + third)", callCount) + if got := callCount.Load(); got != 3 { + t.Errorf("backend called %d times; want 3 (first + second + third)", got) } } @@ -366,7 +377,7 @@ func TestInterrupt_Running(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -408,13 +419,13 @@ func TestSubmit_WhileRunning(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "first", func(Event) {}) + c.Submit(context.Background(), "first") }() waitState(t, c, "thinking") // Concurrent Submit must not start a new turn — always goes to FIFO. - c.Submit(context.Background(), "concurrent", func(Event) {}) + c.Submit(context.Background(), "concurrent") if stored := c.pendingInject.Load(); stored != nil { t.Errorf("concurrent Submit set pendingInject %q; want FIFO only", *stored) @@ -542,11 +553,13 @@ func TestSubmit_PanicRecovery(t *testing.T) { c := newCore(t, be, nil) var errEvents []Event - c.Submit(context.Background(), "hello", func(ev Event) { + sub := c.Bus().Subscribe("event", func(ev Event) { if ev.Role == "error" { errEvents = append(errEvents, ev) } }) + c.Submit(context.Background(), "hello") + c.Bus().Unsubscribe(sub) if len(errEvents) == 0 { t.Error("no error event after panic; expected one") @@ -675,21 +688,40 @@ func TestCommand_I_Empty(t *testing.T) { } func TestCommand_I_SetsInject(t *testing.T) { - c := newCore(t, nil, nil) - collectEvents(context.Background(), c, "/i my inject") + // When running, /i sets pendingInject. + be := defaultBE() + done := make(chan struct{}) + be.respond = func(_ context.Context, msgs []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { + <-done + return textStream("ok"), nil + } + c := newCore(t, be, nil) + go c.Submit(context.Background(), "hello") + waitState(t, c, "thinking") + c.handleCommand(context.Background(), "/i my inject") p := c.pendingInject.Load() if p == nil || *p != "my inject" { t.Errorf("pendingInject = %v; want %q", p, "my inject") } + close(done) } -func TestCommand_I_FallsToFIFOWhenFull(t *testing.T) { - c := newCore(t, nil, nil) - existing := "first" - c.pendingInject.Store(&existing) - collectEvents(context.Background(), c, "/i second") - if got, ok := c.PopQueue(); !ok || got != "second" { - t.Errorf("FIFO after /i with full inject = %q, %v; want %q, true", got, ok, "second") +func TestCommand_I_SubmitsWhenIdle(t *testing.T) { + // When idle, /i submits the prompt directly. + var callCount atomic.Int32 + be := defaultBE() + be.respond = func(_ context.Context, _ []backend.Message, _ []backend.Tool, _ backend.GenerationParams) (<-chan backend.StreamEvent, error) { + callCount.Add(1) + return textStream("ok"), nil + } + c := newCore(t, be, nil) + c.Submit(context.Background(), "/i my prompt") + deadline := time.Now().Add(2 * time.Second) + for callCount.Load() < 1 && time.Now().Before(deadline) { + time.Sleep(5 * time.Millisecond) + } + if callCount.Load() != 1 { + t.Errorf("expected 1 backend call from /i submit; got %d", callCount.Load()) } } @@ -731,7 +763,7 @@ func TestCommand_Compact_WhileRunning(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -799,7 +831,7 @@ func TestCommand_Clear_WhileRunning(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -1021,7 +1053,7 @@ func TestCommand_Backend_WhileRunning(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -1077,7 +1109,7 @@ func TestCommand_Model_WhileRunning(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -1159,7 +1191,7 @@ func TestCommand_Agent_WhileRunning(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -1448,11 +1480,13 @@ func TestRun_RateLimitRetry(t *testing.T) { } c := newCore(t, be, nil) var retryEvents []Event - c.Submit(context.Background(), "hello", func(ev Event) { + sub := c.Bus().Subscribe("event", func(ev Event) { if ev.Role == "retry" { retryEvents = append(retryEvents, ev) } }) + c.Submit(context.Background(), "hello") + c.Bus().Unsubscribe(sub) if callCount != 2 { t.Errorf("backend called %d times; want 2 (retry + success)", callCount) } @@ -1481,7 +1515,7 @@ func TestRun_ToolCancelledBeforeExec(t *testing.T) { return ch, nil } c := newCore(t, be, nil) - c.Submit(ctx, "hello", func(Event) {}) + c.Submit(ctx, "hello") if callCount != 1 { t.Errorf("backend called %d times; want 1 (no retry after cancellation)", callCount) } @@ -1501,7 +1535,7 @@ func TestRun_RequestCancelled(t *testing.T) { return nil, fmt.Errorf("request failed") } c := newCore(t, be, nil) - c.Submit(ctx, "hello", func(Event) {}) + c.Submit(ctx, "hello") if got := c.State(); got != "idle" { t.Errorf("State() = %q after request cancellation; want idle", got) } @@ -1532,7 +1566,7 @@ func TestRun_ExecCancelledWithInject(t *testing.T) { cancel() // ctx cancelled during exec return "", fmt.Errorf("exec cancelled") } - c.Submit(ctx, "hello", func(Event) {}) + c.Submit(ctx, "hello") if callCount != 1 { t.Errorf("backend called %d times; want 1", callCount) } @@ -1629,7 +1663,7 @@ func TestCore_SetGenerationParams_WhileRunning(t *testing.T) { done := make(chan struct{}) go func() { defer close(done) - c.Submit(context.Background(), "hello", func(Event) {}) + c.Submit(context.Background(), "hello") }() waitState(t, c, "thinking") @@ -1802,7 +1836,7 @@ func TestRun_RateLimitRetry_Cancelled(t *testing.T) { return nil, &backend.RateLimitError{RetryAfter: 10 * time.Second} } c := newCore(t, be, nil) - c.Submit(ctx, "hello", func(Event) {}) + c.Submit(ctx, "hello") if callCount != 1 { t.Errorf("backend called %d times after cancelled retry; want 1", callCount) } diff --git a/pkg/agent/types.go b/pkg/agent/types.go index 250f744..a4d27da 100644 --- a/pkg/agent/types.go +++ b/pkg/agent/types.go @@ -5,6 +5,7 @@ import ( "errors" "sync" + "github.com/simonfxr/pubsub" "ollie/pkg/backend" ) @@ -31,13 +32,13 @@ type Event struct { type EventHandler func(Event) // Core is the interface between a frontend (TUI, HTTP handler, etc.) and the -// agent engine. All output from the agent is delivered via EventHandler. +// agent engine. All output from the agent is delivered via the event bus. type Core interface { // Submit processes one line of user input. Slash commands and shell // shortcuts are dispatched synchronously; any other input starts an agent - // turn that streams events to handler until the turn is complete. + // turn that publishes events to the bus until the turn is complete. // After the turn, any queued prompts are drained sequentially. - Submit(ctx context.Context, input string, handler EventHandler) + Submit(ctx context.Context, input string) // Interrupt cancels the current in-progress agent turn. // Returns true if an action was running and was cancelled. @@ -51,6 +52,9 @@ type Core interface { // turn completes. Queue(prompt string) + // Bus returns the session event bus. + Bus() *pubsub.Bus + // PopQueue removes and returns the next queued prompt. // Returns ("", false) if the queue is empty. PopQueue() (string, bool)