elevate: wire turn rate limiter reset on prompt submission

Adds resetElevation callback to sessionHelper, called when a new
prompt is written (turn boundary). The Manager passes
broker.ResetTurn(sessionID) via ManagerConfig.ResetElevation.
This commit is contained in:
Levi Neely 2026-07-29 16:22:28 +02:00
parent c19c8d0cb5
commit 563c6dce2f
4 changed files with 33 additions and 12 deletions

View File

@ -175,6 +175,9 @@ func runServer(sockPath string) {
// D-Bus adapter (initialized after manager so callbacks can reference it). // D-Bus adapter (initialized after manager so callbacks can reference it).
var dbusAdapter *DBusAdapter var dbusAdapter *DBusAdapter
// Elevate broker (initialized after manager; closures capture the pointer).
var elevateBroker *elevate.Broker
mgr := session.NewManager(session.ManagerConfig{ mgr := session.NewManager(session.ManagerConfig{
ToolRegistry: toolRegistry, ToolRegistry: toolRegistry,
SkillsRegistry: skillsRegistry, SkillsRegistry: skillsRegistry,
@ -188,6 +191,11 @@ func runServer(sockPath string) {
Enable9P: !*no9p, Enable9P: !*no9p,
EnableDBus: !*nodbus && (*no9p || *tcpAddr == ""), EnableDBus: !*nodbus && (*no9p || *tcpAddr == ""),
InvalidateModels: modelCache.Invalidate, InvalidateModels: modelCache.Invalidate,
ResetElevation: func(sessionID string) {
if elevateBroker != nil {
elevateBroker.ResetTurn(sessionID)
}
},
OnSessionCreated: func(id string, sess *session.Session) { OnSessionCreated: func(id string, sess *session.Session) {
if dbusAdapter != nil { if dbusAdapter != nil {
dbusAdapter.OnSessionCreated(id, sess) dbusAdapter.OnSessionCreated(id, sess)
@ -233,7 +241,6 @@ func runServer(sockPath string) {
} }
policyPath := filepath.Join(paths.DataDir(), "elevate-policy.yaml") policyPath := filepath.Join(paths.DataDir(), "elevate-policy.yaml")
var elevateBroker *elevate.Broker
{ {
notifyFn := func(req *elevate.Request) { notifyFn := func(req *elevate.Request) {
if dbusAdapter != nil { if dbusAdapter != nil {

View File

@ -210,6 +210,7 @@ func newTestSessionFileStore(t *testing.T, sess *session.Session) (*fs.Tree, *st
func(data []byte) error { return nil }, func(data []byte) error { return nil },
func() {}, func() {},
nil, nil,
nil,
) )
return sf, core return sf, core
} }
@ -217,7 +218,7 @@ func newTestSessionFileStore(t *testing.T, sess *session.Session) (*fs.Tree, *st
func newTestSessionFileStoreWith(t *testing.T, sess *session.Session, kill func(), rename func(string) error, save func([]byte) error) *fs.Tree { func newTestSessionFileStoreWith(t *testing.T, sess *session.Session, kill func(), rename func(string) error, save func([]byte) error) *fs.Tree {
t.Helper() t.Helper()
sink := testSink() sink := testSink()
return session.NewSessionTree(sess, sink.NewLogger("test"), kill, rename, save, func() {}, nil) return session.NewSessionTree(sess, sink.NewLogger("test"), kill, rename, save, func() {}, nil, nil)
} }
// ===== session.Session ===== // ===== session.Session =====
@ -436,7 +437,7 @@ func TestSessionFileStoreReadableContract(t *testing.T) {
defer sess.Cancel() defer sess.Cancel()
sink := testSink() sink := testSink()
_ = session.NewSessionTree(sess, sink.NewLogger("test"), _ = session.NewSessionTree(sess, sink.NewLogger("test"),
func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil) func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil, nil)
} }
func TestSessionFileStoreList(t *testing.T) { func TestSessionFileStoreList(t *testing.T) {
@ -444,7 +445,7 @@ func TestSessionFileStoreList(t *testing.T) {
defer sess.Cancel() defer sess.Cancel()
sink := testSink() sink := testSink()
sf := session.NewSessionTree(sess, sink.NewLogger("test"), sf := session.NewSessionTree(sess, sink.NewLogger("test"),
func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil) func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil, nil)
entries, err := sf.List() entries, err := sf.List()
if err != nil { if err != nil {
@ -461,7 +462,7 @@ func TestSessionFileStoreStatChat(t *testing.T) {
sess.AppendLog([]byte("hello")) sess.AppendLog([]byte("hello"))
sink := testSink() sink := testSink()
sf := session.NewSessionTree(sess, sink.NewLogger("test"), sf := session.NewSessionTree(sess, sink.NewLogger("test"),
func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil) func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil, nil)
fi, err := sf.Stat("chat") fi, err := sf.Stat("chat")
if err != nil { if err != nil {
@ -478,7 +479,7 @@ func TestSessionFileStoreGetChat(t *testing.T) {
sess.AppendLog([]byte("hello")) sess.AppendLog([]byte("hello"))
sink := testSink() sink := testSink()
sf := session.NewSessionTree(sess, sink.NewLogger("test"), sf := session.NewSessionTree(sess, sink.NewLogger("test"),
func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil) func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil, nil)
data := testStoreRead(t, sf, "chat") data := testStoreRead(t, sf, "chat")
if string(data) != "hello" { if string(data) != "hello" {
@ -491,7 +492,7 @@ func TestSessionFileStoreGetContent(t *testing.T) {
defer sess.Cancel() defer sess.Cancel()
sink := testSink() sink := testSink()
sf := session.NewSessionTree(sess, sink.NewLogger("test"), sf := session.NewSessionTree(sess, sink.NewLogger("test"),
func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil) func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil, nil)
for _, name := range []string{"cfg", "offset", "usage", "ctxsz", "models", "systemprompt"} { for _, name := range []string{"cfg", "offset", "usage", "ctxsz", "models", "systemprompt"} {
if _, err := sf.Open(name); err != nil { if _, err := sf.Open(name); err != nil {
@ -505,7 +506,7 @@ func TestSessionFileStorePutCwd(t *testing.T) {
defer sess.Cancel() defer sess.Cancel()
sink := testSink() sink := testSink()
sf := session.NewSessionTree(sess, sink.NewLogger("test"), sf := session.NewSessionTree(sess, sink.NewLogger("test"),
func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil) func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil, nil)
testStoreWrite(t, sf, "cfg", []byte("cwd=/new/path")) testStoreWrite(t, sf, "cfg", []byte("cwd=/new/path"))
core := sess.Core.(*stubCore) core := sess.Core.(*stubCore)
@ -519,7 +520,7 @@ func TestSessionFileStorePutEmpty(t *testing.T) {
defer sess.Cancel() defer sess.Cancel()
sink := testSink() sink := testSink()
sf := session.NewSessionTree(sess, sink.NewLogger("test"), sf := session.NewSessionTree(sess, sink.NewLogger("test"),
func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil) func() {}, func(string) error { return nil }, func([]byte) error { return nil }, func() {}, nil, nil)
// Empty write is a no-op // Empty write is a no-op
e, err := sf.Open("cfg") e, err := sf.Open("cfg")
@ -967,7 +968,7 @@ func TestSessionReactFile(t *testing.T) {
defer cancel() defer cancel()
core := &stubCore{state: "idle"} core := &stubCore{state: "idle"}
sess := session.NewSession("s1", core, ctx, cancel) sess := session.NewSession("s1", core, ctx, cancel)
store := session.NewSessionTree(sess, testSink().NewLogger("test"), func() {}, func(string) error { return nil }, nil, nil, nil) store := session.NewSessionTree(sess, testSink().NewLogger("test"), func() {}, func(string) error { return nil }, nil, nil, nil, nil)
testStoreWrite(t, store, "react", []byte("👍")) testStoreWrite(t, store, "react", []byte("👍"))
if core.reactResponseID != "" || core.reactEmoji != "👍" { if core.reactResponseID != "" || core.reactEmoji != "👍" {

View File

@ -251,6 +251,8 @@ type ManagerConfig struct {
EnableDBus bool EnableDBus bool
// InvalidateModels clears the model cache, forcing a refresh. // InvalidateModels clears the model cache, forcing a refresh.
InvalidateModels func() InvalidateModels func()
// ResetElevation resets the per-turn elevation rate limiter for a session.
ResetElevation func(sessionID string)
// ToolRegistry is the shared tool registry for lazy tool promotion. // ToolRegistry is the shared tool registry for lazy tool promotion.
ToolRegistry *execute.Registry ToolRegistry *execute.Registry
// SkillsRegistry is the shared skills registry for skill loading. // SkillsRegistry is the shared skills registry for skill loading.
@ -706,6 +708,11 @@ func (s *Manager) OpenStore(id string) (*fs.Tree, error) {
} }
func (s *Manager) openStore(sess *Session) (*fs.Tree, error) { func (s *Manager) openStore(sess *Session) (*fs.Tree, error) {
var resetElev func()
if s.cfg.ResetElevation != nil {
id := sess.id
resetElev = func() { s.cfg.ResetElevation(id) }
}
return NewSessionTree( return NewSessionTree(
sess, sess,
s.cfg.Log, s.cfg.Log,
@ -713,6 +720,7 @@ func (s *Manager) openStore(sess *Session) (*fs.Tree, error) {
func(newID string) error { return s.renameSession(sess.id, newID) }, func(newID string) error { return s.renameSession(sess.id, newID) },
nil, // no transcript saving nil, // no transcript saving
s.cfg.InvalidateModels, s.cfg.InvalidateModels,
resetElev,
s.cfg.ToolRegistry, s.cfg.ToolRegistry,
), nil ), nil
} }

View File

@ -53,8 +53,8 @@ func sessionFilePerms() map[string]os.FileMode {
return fs.Perms[fs.PathSessionFile].Files return fs.Perms[fs.PathSessionFile].Files
} }
func NewSessionTree(sess *Session, log *olog.Logger, kill func(), rename func(newID string) error, saveTranscript func([]byte) error, invalidateModels func(), toolRegistry *execute.Registry) *fs.Tree { func NewSessionTree(sess *Session, log *olog.Logger, kill func(), rename func(newID string) error, saveTranscript func([]byte) error, invalidateModels func(), resetElevation func(), toolRegistry *execute.Registry) *fs.Tree {
h := &sessionHelper{sess: sess, log: log, kill: kill, rename: rename, saveTranscript: saveTranscript, invalidateModels: invalidateModels, toolRegistry: toolRegistry} h := &sessionHelper{sess: sess, log: log, kill: kill, rename: rename, saveTranscript: saveTranscript, invalidateModels: invalidateModels, resetElevation: resetElevation, toolRegistry: toolRegistry}
perms := sessionFilePerms() perms := sessionFilePerms()
specs := make([]fs.FileSpec, len(FileList)) specs := make([]fs.FileSpec, len(FileList))
for i, f := range FileList { for i, f := range FileList {
@ -73,6 +73,7 @@ type sessionHelper struct {
rename func(newID string) error rename func(newID string) error
saveTranscript func([]byte) error saveTranscript func([]byte) error
invalidateModels func() invalidateModels func()
resetElevation func() // called on prompt write to reset turn counter
toolRegistry *execute.Registry toolRegistry *execute.Registry
} }
@ -171,6 +172,10 @@ func (h *sessionHelper) fileSpec(name string, mode os.FileMode) fs.FileSpec {
h.sess.InvalidateModelsCache() h.sess.InvalidateModelsCache()
return nil return nil
} }
// Reset elevation rate limiter on new turn
if h.resetElevation != nil {
h.resetElevation()
}
h.sess.mu.Lock() h.sess.mu.Lock()
h.sess.prevPrompt = []byte(input) h.sess.prevPrompt = []byte(input)
h.sess.mu.Unlock() h.sess.mu.Unlock()