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:
parent
c19c8d0cb5
commit
563c6dce2f
9
main.go
9
main.go
|
|
@ -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 {
|
||||||
|
|
|
||||||
19
main_test.go
19
main_test.go
|
|
@ -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 != "👍" {
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
Reference in New Issue