diff --git a/elevate/broker.go b/elevate/broker.go index 84b22d7..64eb4eb 100644 --- a/elevate/broker.go +++ b/elevate/broker.go @@ -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")) diff --git a/elevate/broker_test.go b/elevate/broker_test.go index 9d214f0..2bedf81 100644 --- a/elevate/broker_test.go +++ b/elevate/broker_test.go @@ -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") } } diff --git a/go.mod b/go.mod index dee2ccc..f7f3978 100644 --- a/go.mod +++ b/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 ) diff --git a/go.sum b/go.sum index 7c6dd75..1afcefd 100644 --- a/go.sum +++ b/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= diff --git a/toolsrv/shell.go b/toolsrv/shell.go index b9ae2b4..bb7920d 100644 --- a/toolsrv/shell.go +++ b/toolsrv/shell.go @@ -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()