session: move state/reply to Agent

Agent now owns its execution state (state, reply, stateMu, changeMu,
changeCond). Session delegates State()/Reply()/WaitChange() to Agent.

This is the correct ownership: each agent has independent execution
state. In multi-agent, agents can be idle/thinking independently.
Session remains the coordination layer.
This commit is contained in:
Levi Neely 2026-07-29 20:10:42 +02:00
parent 63f3f494c2
commit 8acadaaa30
4 changed files with 131 additions and 49 deletions

View File

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

View File

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

View File

@ -313,8 +313,6 @@ type Session struct {
id string
uname string
cwd string
state string // "idle", "thinking", "calling: <tool>"
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()

View File

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