elevate: fix session ID passing, leaky bucket rate limit
- executeElevated: use e.sessionID instead of os.Getenv (process-global env is empty in multi-session olliesrv) - executeElevated: merge envExtra into broker env map so session-scoped vars are visible to elevated commands - Replace per-turn rate limit with x/time/rate leaky bucket: burst of 3, refills 1 token per 10s. Only denials/timeouts consume tokens; approvals are free.
This commit is contained in:
parent
517d18d405
commit
92619d4072
|
|
@ -11,13 +11,16 @@ import (
|
|||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"golang.org/x/time/rate"
|
||||
)
|
||||
|
||||
const (
|
||||
RequestTTL = 300 * time.Second
|
||||
FrameData = 'd'
|
||||
FrameExit = 'x'
|
||||
MaxPerTurn = 3 // max elevation requests per agent turn
|
||||
RequestTTL = 300 * time.Second
|
||||
FrameData = 'd'
|
||||
FrameExit = 'x'
|
||||
BurstLimit = 3 // max denied/unanswered requests in burst window
|
||||
RefillInterval = 10 * time.Second // one slot recovers every 10s (30s to full)
|
||||
)
|
||||
|
||||
// NotifyFunc is called when a new request needs human attention.
|
||||
|
|
@ -35,8 +38,8 @@ type Broker struct {
|
|||
pending map[string]*Request // id -> request
|
||||
sessions map[string]*Policy // sessionID -> session policy
|
||||
|
||||
turnMu sync.Mutex
|
||||
turnCount map[string]int // sessionID -> requests this turn
|
||||
rateMu sync.Mutex
|
||||
limiters map[string]*rate.Limiter // sessionID -> denial rate limiter
|
||||
|
||||
listener net.Listener
|
||||
}
|
||||
|
|
@ -80,7 +83,7 @@ func NewBroker(cfg BrokerConfig) (*Broker, error) {
|
|||
sessionValid: cfg.SessionValid,
|
||||
pending: make(map[string]*Request),
|
||||
sessions: make(map[string]*Policy),
|
||||
turnCount: make(map[string]int),
|
||||
limiters: make(map[string]*rate.Limiter),
|
||||
listener: ln,
|
||||
}
|
||||
|
||||
|
|
@ -169,34 +172,38 @@ func (b *Broker) RemoveSession(sessionID string) {
|
|||
defer b.mu.Unlock()
|
||||
delete(b.sessions, sessionID)
|
||||
|
||||
b.turnMu.Lock()
|
||||
delete(b.turnCount, sessionID)
|
||||
b.turnMu.Unlock()
|
||||
b.rateMu.Lock()
|
||||
delete(b.limiters, sessionID)
|
||||
b.rateMu.Unlock()
|
||||
}
|
||||
|
||||
// ResetTurn resets the per-turn elevation counter for a session.
|
||||
// Call this at the start of each new user turn (prompt submission).
|
||||
func (b *Broker) ResetTurn(sessionID string) {
|
||||
// limiterFor returns (or creates) the denial rate limiter for a session.
|
||||
// The limiter starts full (BurstLimit tokens available). Tokens are only
|
||||
// consumed on denial/timeout, not on approval. One token refills every
|
||||
// RefillInterval (10s by default), so a full burst window is 30s.
|
||||
func (b *Broker) limiterFor(sessionID string) *rate.Limiter {
|
||||
b.rateMu.Lock()
|
||||
defer b.rateMu.Unlock()
|
||||
lim, ok := b.limiters[sessionID]
|
||||
if !ok {
|
||||
lim = rate.NewLimiter(rate.Every(RefillInterval), BurstLimit)
|
||||
b.limiters[sessionID] = lim
|
||||
}
|
||||
return lim
|
||||
}
|
||||
|
||||
// checkRateLimit returns true if the session still has denial budget.
|
||||
func (b *Broker) checkRateLimit(sessionID string) bool {
|
||||
return b.limiterFor(sessionID).Tokens() >= 1
|
||||
}
|
||||
|
||||
// recordDenial consumes a token from the session's denial budget.
|
||||
// Called when a request is denied or times out (user not responding).
|
||||
func (b *Broker) recordDenial(sessionID string) {
|
||||
if sessionID == "" {
|
||||
return
|
||||
}
|
||||
b.turnMu.Lock()
|
||||
delete(b.turnCount, sessionID)
|
||||
b.turnMu.Unlock()
|
||||
}
|
||||
|
||||
// checkTurnLimit increments and checks the per-turn counter for a session.
|
||||
// Returns true if the request is allowed, false if rate-limited.
|
||||
// Callers must ensure sessionID is non-empty before calling.
|
||||
func (b *Broker) checkTurnLimit(sessionID string) bool {
|
||||
b.turnMu.Lock()
|
||||
defer b.turnMu.Unlock()
|
||||
n := b.turnCount[sessionID]
|
||||
if n >= MaxPerTurn {
|
||||
return false
|
||||
}
|
||||
b.turnCount[sessionID] = n + 1
|
||||
return true
|
||||
b.limiterFor(sessionID).Allow()
|
||||
}
|
||||
|
||||
// GlobalPolicy returns the policy store for 9P access.
|
||||
|
|
@ -283,10 +290,10 @@ func (b *Broker) handleConn(conn net.Conn) {
|
|||
return
|
||||
}
|
||||
|
||||
// Rate limit: max 3 elevation requests per agent turn
|
||||
if !b.checkTurnLimit(sessionID) {
|
||||
// Rate limit: leaky bucket — 3 denied/unanswered requests per 30s window
|
||||
if !b.checkRateLimit(sessionID) {
|
||||
b.logf("elevate: rate-limited session=%s cmd=%q", sessionID, msg.Cmd)
|
||||
sendFrame(conn, FrameData, []byte("elevate-broker: rate-limited (max 3 per turn)\n"))
|
||||
sendFrame(conn, FrameData, []byte("elevate-broker: rate-limited (too many denied/unanswered requests)\n"))
|
||||
sendFrame(conn, FrameExit, []byte("1"))
|
||||
return
|
||||
}
|
||||
|
|
@ -339,6 +346,7 @@ func (b *Broker) handleConn(conn net.Conn) {
|
|||
b.logf("elevate: approved id=%s cmd=%q", req.ID, msg.Cmd)
|
||||
b.executeAndStream(conn, msg.Cmd, msg.Cwd, msg.Env)
|
||||
default:
|
||||
b.recordDenial(sessionID)
|
||||
b.logf("elevate: denied id=%s cmd=%q reason=%s", req.ID, msg.Cmd, res)
|
||||
sendFrame(conn, FrameData, []byte(fmt.Sprintf("elevate-broker: request %s by user\n", res)))
|
||||
sendFrame(conn, FrameExit, []byte("1"))
|
||||
|
|
|
|||
|
|
@ -236,72 +236,73 @@ func TestBrokerPersist(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestBrokerRateLimitPerTurn(t *testing.T) {
|
||||
func TestBrokerRateLimitLeakyBucket(t *testing.T) {
|
||||
b, _ := startTestBroker(t)
|
||||
|
||||
sessionID := "test-session"
|
||||
|
||||
// Simulate 3 requests (all should be allowed)
|
||||
for i := 0; i < MaxPerTurn; i++ {
|
||||
if !b.checkTurnLimit(sessionID) {
|
||||
t.Fatalf("request %d should not be rate-limited", i+1)
|
||||
// No denials — should be allowed
|
||||
for i := 0; i < BurstLimit+1; i++ {
|
||||
if !b.checkRateLimit(sessionID) {
|
||||
t.Fatalf("request %d should not be rate-limited (no denials yet)", i+1)
|
||||
}
|
||||
}
|
||||
|
||||
// 4th request should be rate-limited
|
||||
if b.checkTurnLimit(sessionID) {
|
||||
t.Error("4th request should be rate-limited")
|
||||
// Record BurstLimit denials — exhaust the bucket
|
||||
for i := 0; i < BurstLimit; i++ {
|
||||
b.recordDenial(sessionID)
|
||||
}
|
||||
|
||||
// Reset turn counter
|
||||
b.ResetTurn(sessionID)
|
||||
|
||||
// After reset, requests should be allowed again
|
||||
if !b.checkTurnLimit(sessionID) {
|
||||
t.Error("request after ResetTurn should be allowed")
|
||||
// Now should be rate-limited
|
||||
if b.checkRateLimit(sessionID) {
|
||||
t.Error("should be rate-limited after BurstLimit denials")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerRateLimitDenialMessage(t *testing.T) {
|
||||
func TestBrokerRateLimitRefill(t *testing.T) {
|
||||
b, _ := startTestBroker(t)
|
||||
|
||||
// Test the unit logic directly:
|
||||
// identifySession returns "" for test connections (no PEERCRED mapping),
|
||||
// and checkTurnLimit allows unknown sessions. Rate limiting only applies
|
||||
// when session identity is established.
|
||||
sid := "rate-test"
|
||||
for i := 0; i < MaxPerTurn; i++ {
|
||||
if !b.checkTurnLimit(sid) {
|
||||
t.Fatalf("request %d should be allowed", i+1)
|
||||
}
|
||||
sid := "refill-test"
|
||||
|
||||
// Exhaust the bucket
|
||||
for i := 0; i < BurstLimit; i++ {
|
||||
b.recordDenial(sid)
|
||||
}
|
||||
if b.checkTurnLimit(sid) {
|
||||
t.Fatal("should be rate-limited after MaxPerTurn")
|
||||
if b.checkRateLimit(sid) {
|
||||
t.Fatal("should be rate-limited")
|
||||
}
|
||||
|
||||
// Verify the counter survives and then resets
|
||||
b.ResetTurn(sid)
|
||||
if !b.checkTurnLimit(sid) {
|
||||
t.Fatal("should be allowed after reset")
|
||||
// Manually set the limiter to allow (simulates time passing)
|
||||
// We can't easily wait 10s in a unit test, so verify RemoveSession resets.
|
||||
b.RemoveSession(sid)
|
||||
if !b.checkRateLimit(sid) {
|
||||
t.Fatal("should be allowed after RemoveSession (fresh limiter)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBrokerRemoveSessionCleansUpTurnCount(t *testing.T) {
|
||||
func TestBrokerRemoveSessionCleansUpLimiter(t *testing.T) {
|
||||
b, _ := startTestBroker(t)
|
||||
|
||||
sid := "cleanup-test"
|
||||
|
||||
// Use up some turn budget
|
||||
b.checkTurnLimit(sid)
|
||||
b.checkTurnLimit(sid)
|
||||
// Record some denials
|
||||
b.recordDenial(sid)
|
||||
b.recordDenial(sid)
|
||||
|
||||
// Remove session
|
||||
b.RemoveSession(sid)
|
||||
|
||||
// Counter should be gone — new requests allowed
|
||||
for i := 0; i < MaxPerTurn; i++ {
|
||||
if !b.checkTurnLimit(sid) {
|
||||
t.Fatalf("request %d after RemoveSession should be allowed", i+1)
|
||||
}
|
||||
// Fresh limiter — exhaust again
|
||||
for i := 0; i < BurstLimit; i++ {
|
||||
b.recordDenial(sid)
|
||||
}
|
||||
if b.checkRateLimit(sid) {
|
||||
t.Fatal("should be rate-limited after BurstLimit denials")
|
||||
}
|
||||
|
||||
// Clean up and verify fresh state
|
||||
b.RemoveSession(sid)
|
||||
if !b.checkRateLimit(sid) {
|
||||
t.Fatal("should be allowed after RemoveSession")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
1
go.mod
1
go.mod
|
|
@ -11,4 +11,5 @@ require github.com/simonfxr/pubsub v0.0.5
|
|||
require (
|
||||
golang.org/x/crypto v0.54.0 // indirect
|
||||
golang.org/x/sys v0.47.0 // indirect
|
||||
golang.org/x/time v0.15.0 // indirect
|
||||
)
|
||||
|
|
|
|||
2
go.sum
2
go.sum
|
|
@ -10,6 +10,8 @@ golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw=
|
|||
golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
|
||||
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
|
|
|
|||
|
|
@ -142,19 +142,24 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
|
|||
return "", fmt.Errorf("elevation not available: %w", err)
|
||||
}
|
||||
|
||||
// Send request
|
||||
// Send request — merge process env with session-scoped envExtra.
|
||||
envMap := make(map[string]string, len(os.Environ()))
|
||||
for _, kv := range os.Environ() {
|
||||
if k, v, ok := strings.Cut(kv, "="); ok {
|
||||
envMap[k] = v
|
||||
}
|
||||
}
|
||||
e.envMu.RLock()
|
||||
for k, v := range e.envExtra {
|
||||
envMap[k] = v
|
||||
}
|
||||
e.envMu.RUnlock()
|
||||
reqJSON, _ := json.Marshal(struct {
|
||||
Cmd string `json:"cmd"`
|
||||
Cwd string `json:"cwd"`
|
||||
Env map[string]string `json:"env"`
|
||||
Session string `json:"session,omitempty"`
|
||||
}{Cmd: cmd, Cwd: dir, Env: envMap, Session: os.Getenv("OLLIE_SESSION_ID")})
|
||||
}{Cmd: cmd, Cwd: dir, Env: envMap, Session: e.sessionID})
|
||||
reqJSON = append(reqJSON, '\n')
|
||||
if _, err := conn.Write(reqJSON); err != nil {
|
||||
conn.Close()
|
||||
|
|
|
|||
Reference in New Issue