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:
parent
63f3f494c2
commit
8acadaaa30
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
Reference in New Issue