diff --git a/session/agent.go b/session/agent.go index 8a8fa02..6e5b01f 100644 --- a/session/agent.go +++ b/session/agent.go @@ -1,6 +1,7 @@ package session import ( + "context" "sync" "sync/atomic" @@ -25,6 +26,13 @@ type Agent struct { currentAction atomic.Pointer[actionHandle] warnedContext bool resultCache sync.Map + + // Execution state — owned by the agent, protected by stateMu. + state string // "idle", "thinking", "calling: " + reply string // last assistant response + stateMu sync.RWMutex + changeMu sync.Mutex + changeCond *sync.Cond } // Backend returns the active backend from the runtime. @@ -53,3 +61,83 @@ func (ag *Agent) ModelName() string { } 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) +} diff --git a/session/commands.go b/session/commands.go index adc3365..cd75347 100644 --- a/session/commands.go +++ b/session/commands.go @@ -198,7 +198,7 @@ func (a *Session) handleCommand(ctx context.Context, input string) bool { a.r.agentName = name a.r.history = nil a.pushSessionEnv() - a.notifyChange() + a.r.notifyChange() for _, msg := range rt.Messages { a.emit(infoEvent(msg)) } diff --git a/session/session.go b/session/session.go index ce5df95..ae43969 100644 --- a/session/session.go +++ b/session/session.go @@ -313,8 +313,6 @@ type Session struct { id string uname string cwd string - state string // "idle", "thinking", "calling: " - reply string bus *pubsub.Bus fifo Fifo envMu sync.RWMutex @@ -335,8 +333,6 @@ type Session struct { toolCallCount atomic.Int64 pendingInject atomic.Pointer[string] mu sync.RWMutex - changeMu sync.Mutex - changeCond *sync.Cond submitMu sync.Mutex // serializes Submit calls (commands + turns) auditLog *olog.Logger @@ -471,7 +467,6 @@ func New(cfg Config) *Session { id: cfg.SessionID, uname: cfg.Uname, cwd: paths.ExpandHome(cfg.CWD), - state: "idle", bus: pubsub.NewBus(), env: make(map[string]string), r: &Agent{ @@ -492,7 +487,8 @@ func New(cfg Config) *Session { readPlanStep: readPlanStep, listHandlers: cfg.ListHandlers, } - a.changeCond = sync.NewCond(&a.changeMu) + a.r.InitCond() + a.r.state = "idle" a.turnError = func(_ context.Context, errType, errMsg string) HookResult { hookCtx, cancel := context.WithTimeout(context.Background(), time.Duration(hookTimeout)*time.Second) defer cancel() @@ -722,28 +718,27 @@ func (a *Session) ModelName() string { func (a *Session) Agent() *Agent { return a.r } func (a *Session) State() string { - return a.state -} - -func (a *Session) notifyChange() { - a.changeMu.Lock() - a.changeCond.Broadcast() - a.changeMu.Unlock() + return a.r.State() } func (a *Session) setState(state string) { - a.state = state a.log.Debug("state -> %q", state) - a.notifyChange() + a.r.SetState(state) } // WaitChange blocks until the named field changes from current, then returns // the new value. Returns ("", false) if ctx is cancelled. func (a *Session) WaitChange(ctx context.Context, field, current string) (string, bool) { + // For state changes, delegate directly to the agent. + if field == WatchState { + return a.r.WaitChange(ctx, field, current) + } + + // Other watched fields (usage, ctxsz, cwd, agent) are session-level; + // they change as a side effect of agent activity, so we still listen + // on the agent's change signal. read := func() string { switch field { - case WatchState: - return a.State() case WatchUsage: return a.Usage() case WatchCtxSz: @@ -756,29 +751,26 @@ func (a *Session) WaitChange(ctx context.Context, field, current string) (string return "" } - // context.AfterFunc fires in a separate goroutine when ctx is done, - // broadcasting to unblock any waiters. stop := context.AfterFunc(ctx, func() { - a.changeMu.Lock() - a.changeCond.Broadcast() - a.changeMu.Unlock() + a.r.changeMu.Lock() + a.r.changeCond.Broadcast() + a.r.changeMu.Unlock() }) defer stop() - a.changeMu.Lock() - defer a.changeMu.Unlock() + a.r.changeMu.Lock() + defer a.r.changeMu.Unlock() for ctx.Err() == nil { if v := read(); v != current { return v, true } - a.changeCond.Wait() + a.r.changeCond.Wait() } return "", false } - func (a *Session) Reply() string { - r := a.reply + r := a.r.Reply() a.log.Debug("Reply() len=%d", len(r)) return r } @@ -818,7 +810,7 @@ func (a *Session) SetCWD(dir string) error { } } } - a.notifyChange() + a.r.notifyChange() return nil } @@ -1191,7 +1183,7 @@ func (a *Session) Submit(ctx context.Context, input string) { if a := a.r.currentAction.Swap(nil); a != nil { a.cancel(fmt.Errorf("%v", r)) } - a.setState("idle") + a.r.SetState("idle") a.emit(Event{Role: "error", Content: fmt.Sprintf("%v", r)}) } }() @@ -1279,7 +1271,7 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { actCtx, actCancel := context.WithCancelCause(ctx) handle := &actionHandle{cancel: actCancel} a.r.currentAction.Store(handle) - a.setState("thinking") + a.r.SetState("thinking") a.auditLog.Debug("turn: start input=%s session=%s", auditTruncate(input), a.id) @@ -1303,14 +1295,14 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { case "assistant": replyBuf.WriteString(ev.Content) case "call": - a.setState("calling: " + ev.Name) + a.r.SetState("calling: " + ev.Name) a.auditLog.Debug("call: %s %s", ev.Name, auditTruncate(string(ev.Content))) case "tool": a.auditLog.Debug("result: %s %s", ev.Name, auditTruncate(ev.Content)) case "state": - a.setState(ev.Content) + a.r.SetState(ev.Content) case "limitretry": - a.setState("limitretry") + a.r.SetState("limitretry") case "error": a.auditLog.Debug("error: %s", ev.Content) } @@ -1325,7 +1317,7 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { OutputTokens: out, CostUSD: costUSD, }, est != 0) - a.notifyChange() + a.r.notifyChange() if maxCostStr := os.Getenv("OLLIE_MAX_SESSION_COST"); maxCostStr != "" { if limit, ferr := strconv.ParseFloat(maxCostStr, 64); ferr == nil && limit > 0 && a.r.history.SessionCostUSD >= limit { a.emit(infoEvent(fmt.Sprintf("spending cap $%.2f reached — stopping", limit))) @@ -1372,11 +1364,11 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { return } a.emit(Event{Role: "info", Content: "auto-compacting context...\n"}) - a.setState("compacting") + a.r.SetState("compacting") if _, err := a.runCompact(ctx, "auto"); err != nil { panic(fmt.Sprintf("mid-turn auto-compact: %v", err)) } - a.setState("thinking") + a.r.SetState("thinking") } // Warn once when context usage crosses 60%; compact at 75%. @@ -1384,11 +1376,11 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { tokens := a.r.history.estimateTokens() if compactLimit := a.autoCompactLimit(ctx); compactLimit > 0 && tokens >= compactLimit { a.emit(Event{Role: "info", Content: "auto-compacting context...\n"}) - a.setState("compacting") + a.r.SetState("compacting") if _, err := a.runCompact(ctx, "auto"); err != nil { panic(fmt.Sprintf("auto-compact: %v", err)) } - a.setState("thinking") + a.r.SetState("thinking") } else if warnLimit := a.autoWarnLimit(ctx); warnLimit > 0 && tokens >= warnLimit && !a.r.warnedContext { ctxLen := a.r.cfg.Backend.ContextLength(ctx) if ctxLen <= 0 { @@ -1405,7 +1397,7 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { if limit, err := strconv.ParseFloat(maxCostStr, 64); err == nil && limit > 0 { if a.r.history.SessionCostUSD >= limit { a.emit(Event{Role: "error", Content: fmt.Sprintf("spending cap $%.2f reached (session total $%.4f)", limit, a.r.history.SessionCostUSD)}) - a.setState("idle") + a.r.SetState("idle") actCancel(nil) a.r.currentAction.CompareAndSwap(handle, nil) if snapSession == nil { @@ -1443,12 +1435,12 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { overflowRetried = true a.r.history.messages = snapMessages a.emit(Event{Role: "info", Content: "context overflow — compacting and retrying...\n"}) - a.setState("compacting") + a.r.SetState("compacting") if _, cerr := a.runCompact(ctx, "overflow"); cerr != nil { break } a.r.history.appendUserMessage(input) - a.setState("thinking") + a.r.SetState("thinking") a.r.history.resetTurnAccumulators() replyBuf.Reset() actCtx, actCancel = context.WithCancelCause(ctx) @@ -1460,10 +1452,10 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { } a.mu.Lock() - a.reply = replyBuf.String() + a.r.SetReply(replyBuf.String()) a.mu.Unlock() replyBuf.Reset() - a.setState("idle") + a.r.SetState("idle") a.flushSave() if err != nil { @@ -1506,8 +1498,8 @@ func (a *Session) executeTurn(ctx context.Context, input string) string { a.emit(Event{Role: "info", Content: fmt.Sprintf("costLast=$%.4f\n", a.r.history.LastTurnCostUSD)}) } a.auditLog.Debug("turn: end reply=%s cost=$%.4f session_total=$%.4f session=%s", - auditTruncate(a.reply), a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD, a.id) - a.notifyChange() + auditTruncate(a.r.Reply()), a.r.history.LastTurnCostUSD, a.r.history.SessionCostUSD, a.id) + a.r.notifyChange() } a.saveSession() diff --git a/session/waitchange_test.go b/session/waitchange_test.go index 702ba2f..3fc5aa3 100644 --- a/session/waitchange_test.go +++ b/session/waitchange_test.go @@ -3,7 +3,6 @@ package session import ( "context" "io" - "sync" "github.com/simonfxr/pubsub" "testing" @@ -14,14 +13,17 @@ import ( // newTestCore returns a minimal Session for testing. func newTestCore(initialState string) *Session { + ag := &Agent{ + state: initialState, + } + ag.InitCond() a := &Session{ id: "test", - state: initialState, bus: pubsub.NewBus(), env: make(map[string]string), + r: ag, log: olog.NewWriter("test", olog.LevelError+1, io.Discard, io.Discard), } - a.changeCond = sync.NewCond(&a.changeMu) return a }