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"
|
"sync"
|
||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"golang.org/x/time/rate"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
RequestTTL = 300 * time.Second
|
RequestTTL = 300 * time.Second
|
||||||
FrameData = 'd'
|
FrameData = 'd'
|
||||||
FrameExit = 'x'
|
FrameExit = 'x'
|
||||||
MaxPerTurn = 3 // max elevation requests per agent turn
|
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.
|
// NotifyFunc is called when a new request needs human attention.
|
||||||
|
|
@ -35,8 +38,8 @@ type Broker struct {
|
||||||
pending map[string]*Request // id -> request
|
pending map[string]*Request // id -> request
|
||||||
sessions map[string]*Policy // sessionID -> session policy
|
sessions map[string]*Policy // sessionID -> session policy
|
||||||
|
|
||||||
turnMu sync.Mutex
|
rateMu sync.Mutex
|
||||||
turnCount map[string]int // sessionID -> requests this turn
|
limiters map[string]*rate.Limiter // sessionID -> denial rate limiter
|
||||||
|
|
||||||
listener net.Listener
|
listener net.Listener
|
||||||
}
|
}
|
||||||
|
|
@ -80,7 +83,7 @@ func NewBroker(cfg BrokerConfig) (*Broker, error) {
|
||||||
sessionValid: cfg.SessionValid,
|
sessionValid: cfg.SessionValid,
|
||||||
pending: make(map[string]*Request),
|
pending: make(map[string]*Request),
|
||||||
sessions: make(map[string]*Policy),
|
sessions: make(map[string]*Policy),
|
||||||
turnCount: make(map[string]int),
|
limiters: make(map[string]*rate.Limiter),
|
||||||
listener: ln,
|
listener: ln,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -169,34 +172,38 @@ func (b *Broker) RemoveSession(sessionID string) {
|
||||||
defer b.mu.Unlock()
|
defer b.mu.Unlock()
|
||||||
delete(b.sessions, sessionID)
|
delete(b.sessions, sessionID)
|
||||||
|
|
||||||
b.turnMu.Lock()
|
b.rateMu.Lock()
|
||||||
delete(b.turnCount, sessionID)
|
delete(b.limiters, sessionID)
|
||||||
b.turnMu.Unlock()
|
b.rateMu.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// ResetTurn resets the per-turn elevation counter for a session.
|
// limiterFor returns (or creates) the denial rate limiter for a session.
|
||||||
// Call this at the start of each new user turn (prompt submission).
|
// The limiter starts full (BurstLimit tokens available). Tokens are only
|
||||||
func (b *Broker) ResetTurn(sessionID string) {
|
// 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 == "" {
|
if sessionID == "" {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
b.turnMu.Lock()
|
b.limiterFor(sessionID).Allow()
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// GlobalPolicy returns the policy store for 9P access.
|
// GlobalPolicy returns the policy store for 9P access.
|
||||||
|
|
@ -283,10 +290,10 @@ func (b *Broker) handleConn(conn net.Conn) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// Rate limit: max 3 elevation requests per agent turn
|
// Rate limit: leaky bucket — 3 denied/unanswered requests per 30s window
|
||||||
if !b.checkTurnLimit(sessionID) {
|
if !b.checkRateLimit(sessionID) {
|
||||||
b.logf("elevate: rate-limited session=%s cmd=%q", sessionID, msg.Cmd)
|
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"))
|
sendFrame(conn, FrameExit, []byte("1"))
|
||||||
return
|
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.logf("elevate: approved id=%s cmd=%q", req.ID, msg.Cmd)
|
||||||
b.executeAndStream(conn, msg.Cmd, msg.Cwd, msg.Env)
|
b.executeAndStream(conn, msg.Cmd, msg.Cwd, msg.Env)
|
||||||
default:
|
default:
|
||||||
|
b.recordDenial(sessionID)
|
||||||
b.logf("elevate: denied id=%s cmd=%q reason=%s", req.ID, msg.Cmd, res)
|
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, FrameData, []byte(fmt.Sprintf("elevate-broker: request %s by user\n", res)))
|
||||||
sendFrame(conn, FrameExit, []byte("1"))
|
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)
|
b, _ := startTestBroker(t)
|
||||||
|
|
||||||
sessionID := "test-session"
|
sessionID := "test-session"
|
||||||
|
|
||||||
// Simulate 3 requests (all should be allowed)
|
// No denials — should be allowed
|
||||||
for i := 0; i < MaxPerTurn; i++ {
|
for i := 0; i < BurstLimit+1; i++ {
|
||||||
if !b.checkTurnLimit(sessionID) {
|
if !b.checkRateLimit(sessionID) {
|
||||||
t.Fatalf("request %d should not be rate-limited", i+1)
|
t.Fatalf("request %d should not be rate-limited (no denials yet)", i+1)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 4th request should be rate-limited
|
// Record BurstLimit denials — exhaust the bucket
|
||||||
if b.checkTurnLimit(sessionID) {
|
for i := 0; i < BurstLimit; i++ {
|
||||||
t.Error("4th request should be rate-limited")
|
b.recordDenial(sessionID)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Reset turn counter
|
// Now should be rate-limited
|
||||||
b.ResetTurn(sessionID)
|
if b.checkRateLimit(sessionID) {
|
||||||
|
t.Error("should be rate-limited after BurstLimit denials")
|
||||||
// After reset, requests should be allowed again
|
|
||||||
if !b.checkTurnLimit(sessionID) {
|
|
||||||
t.Error("request after ResetTurn should be allowed")
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBrokerRateLimitDenialMessage(t *testing.T) {
|
func TestBrokerRateLimitRefill(t *testing.T) {
|
||||||
b, _ := startTestBroker(t)
|
b, _ := startTestBroker(t)
|
||||||
|
|
||||||
// Test the unit logic directly:
|
sid := "refill-test"
|
||||||
// identifySession returns "" for test connections (no PEERCRED mapping),
|
|
||||||
// and checkTurnLimit allows unknown sessions. Rate limiting only applies
|
// Exhaust the bucket
|
||||||
// when session identity is established.
|
for i := 0; i < BurstLimit; i++ {
|
||||||
sid := "rate-test"
|
b.recordDenial(sid)
|
||||||
for i := 0; i < MaxPerTurn; i++ {
|
|
||||||
if !b.checkTurnLimit(sid) {
|
|
||||||
t.Fatalf("request %d should be allowed", i+1)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if b.checkTurnLimit(sid) {
|
if b.checkRateLimit(sid) {
|
||||||
t.Fatal("should be rate-limited after MaxPerTurn")
|
t.Fatal("should be rate-limited")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Verify the counter survives and then resets
|
// Manually set the limiter to allow (simulates time passing)
|
||||||
b.ResetTurn(sid)
|
// We can't easily wait 10s in a unit test, so verify RemoveSession resets.
|
||||||
if !b.checkTurnLimit(sid) {
|
b.RemoveSession(sid)
|
||||||
t.Fatal("should be allowed after reset")
|
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)
|
b, _ := startTestBroker(t)
|
||||||
|
|
||||||
sid := "cleanup-test"
|
sid := "cleanup-test"
|
||||||
|
|
||||||
// Use up some turn budget
|
// Record some denials
|
||||||
b.checkTurnLimit(sid)
|
b.recordDenial(sid)
|
||||||
b.checkTurnLimit(sid)
|
b.recordDenial(sid)
|
||||||
|
|
||||||
// Remove session
|
// Remove session
|
||||||
b.RemoveSession(sid)
|
b.RemoveSession(sid)
|
||||||
|
|
||||||
// Counter should be gone — new requests allowed
|
// Fresh limiter — exhaust again
|
||||||
for i := 0; i < MaxPerTurn; i++ {
|
for i := 0; i < BurstLimit; i++ {
|
||||||
if !b.checkTurnLimit(sid) {
|
b.recordDenial(sid)
|
||||||
t.Fatalf("request %d after RemoveSession should be allowed", i+1)
|
}
|
||||||
}
|
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 (
|
require (
|
||||||
golang.org/x/crypto v0.54.0 // indirect
|
golang.org/x/crypto v0.54.0 // indirect
|
||||||
golang.org/x/sys v0.47.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/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 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
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 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
|
||||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
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=
|
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)
|
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()))
|
envMap := make(map[string]string, len(os.Environ()))
|
||||||
for _, kv := range os.Environ() {
|
for _, kv := range os.Environ() {
|
||||||
if k, v, ok := strings.Cut(kv, "="); ok {
|
if k, v, ok := strings.Cut(kv, "="); ok {
|
||||||
envMap[k] = v
|
envMap[k] = v
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
e.envMu.RLock()
|
||||||
|
for k, v := range e.envExtra {
|
||||||
|
envMap[k] = v
|
||||||
|
}
|
||||||
|
e.envMu.RUnlock()
|
||||||
reqJSON, _ := json.Marshal(struct {
|
reqJSON, _ := json.Marshal(struct {
|
||||||
Cmd string `json:"cmd"`
|
Cmd string `json:"cmd"`
|
||||||
Cwd string `json:"cwd"`
|
Cwd string `json:"cwd"`
|
||||||
Env map[string]string `json:"env"`
|
Env map[string]string `json:"env"`
|
||||||
Session string `json:"session,omitempty"`
|
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')
|
reqJSON = append(reqJSON, '\n')
|
||||||
if _, err := conn.Write(reqJSON); err != nil {
|
if _, err := conn.Write(reqJSON); err != nil {
|
||||||
conn.Close()
|
conn.Close()
|
||||||
|
|
|
||||||
Reference in New Issue