package elevate import ( "encoding/binary" "encoding/json" "fmt" "net" "os" "os/exec" "strings" "sync" "syscall" "time" "golang.org/x/time/rate" ) const ( 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. // Implementations should show a notification or log the request. type NotifyFunc func(req *Request) // Broker manages elevation requests. type Broker struct { policy *PolicyStore notify NotifyFunc logf func(string, ...any) sessionValid func(string) bool mu sync.RWMutex pending map[string]*Request // id -> request sessions map[string]*Policy // sessionID -> session policy rateMu sync.Mutex limiters map[string]*rate.Limiter // sessionID -> denial rate limiter listener net.Listener } // BrokerConfig configures the broker. type BrokerConfig struct { SocketPath string PolicyPath string // disk path for global policy YAML Notify NotifyFunc Logf func(string, ...any) SessionValid func(sessionID string) bool // returns true if session exists } // NewBroker creates and starts the elevation broker. func NewBroker(cfg BrokerConfig) (*Broker, error) { if cfg.Logf == nil { cfg.Logf = func(string, ...any) {} } if cfg.Notify == nil { cfg.Notify = func(*Request) {} } // Ensure socket directory exists dir := socketDir(cfg.SocketPath) os.MkdirAll(dir, 0700) //nolint:errcheck // Remove stale socket os.Remove(cfg.SocketPath) //nolint:errcheck ln, err := net.Listen("unix", cfg.SocketPath) if err != nil { return nil, fmt.Errorf("elevate: listen %s: %w", cfg.SocketPath, err) } // Restrict socket permissions os.Chmod(cfg.SocketPath, 0600) //nolint:errcheck b := &Broker{ policy: NewPolicyStore(cfg.PolicyPath), notify: cfg.Notify, logf: cfg.Logf, sessionValid: cfg.SessionValid, pending: make(map[string]*Request), sessions: make(map[string]*Policy), limiters: make(map[string]*rate.Limiter), listener: ln, } go b.acceptLoop() cfg.Logf("elevate: listening on %s", cfg.SocketPath) return b, nil } // Close shuts down the broker. func (b *Broker) Close() { b.listener.Close() } // Pending returns all pending requests. func (b *Broker) Pending() []*Request { b.mu.RLock() defer b.mu.RUnlock() out := make([]*Request, 0, len(b.pending)) for _, r := range b.pending { out = append(out, r) } return out } // PendingByID returns a specific pending request. func (b *Broker) PendingByID(id string) *Request { b.mu.RLock() defer b.mu.RUnlock() return b.pending[id] } // Resolve resolves a pending request. Safe to call after timeout (no-op). func (b *Broker) Resolve(id string, res Resolution) bool { b.mu.Lock() req, ok := b.pending[id] if !ok { b.mu.Unlock() return false } delete(b.pending, id) b.mu.Unlock() // If persist, add rule. Session-scoped if we know the session, global otherwise. if res == ResolvePersist { if req.SessionID != "" { b.addSessionRule(req.SessionID, Rule{Cmd: req.Cmd}) } else { b.policy.AddGlobal(Rule{Cmd: req.Cmd}) //nolint:errcheck } res = ResolveApprove // treat as approve for execution } select { case req.resolved <- res: default: } return true } // SessionPolicy returns the effective policy for a session (global + session). func (b *Broker) SessionPolicy(sessionID string) *Policy { global := b.policy.Global() b.mu.RLock() sess, ok := b.sessions[sessionID] b.mu.RUnlock() // Merge: global rules + session rules merged := &Policy{Rules: make([]Rule, len(global.Rules))} copy(merged.Rules, global.Rules) if ok { merged.Rules = append(merged.Rules, sess.Rules...) } return merged } // SetSessionPolicy replaces the session policy. func (b *Broker) SetSessionPolicy(sessionID string, p Policy) { b.mu.Lock() defer b.mu.Unlock() b.sessions[sessionID] = &p } // RemoveSession cleans up session policy when a session is killed. func (b *Broker) RemoveSession(sessionID string) { b.mu.Lock() defer b.mu.Unlock() delete(b.sessions, sessionID) b.rateMu.Lock() delete(b.limiters, sessionID) b.rateMu.Unlock() } // 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.limiterFor(sessionID).Allow() } // GlobalPolicy returns the policy store for 9P access. func (b *Broker) GlobalPolicy() *PolicyStore { return b.policy } func (b *Broker) addSessionRule(sessionID string, r Rule) { if sessionID == "" { return } b.mu.Lock() defer b.mu.Unlock() p, ok := b.sessions[sessionID] if !ok { p = &Policy{} b.sessions[sessionID] = p } p.Add(r) } func (b *Broker) acceptLoop() { for { conn, err := b.listener.Accept() if err != nil { return // listener closed } go b.handleConn(conn) } } func (b *Broker) handleConn(conn net.Conn) { defer conn.Close() // Read request (JSON terminated by newline) conn.SetReadDeadline(time.Now().Add(10 * time.Second)) //nolint:errcheck buf := make([]byte, 0, 4096) tmp := make([]byte, 4096) for { n, err := conn.Read(tmp) if n > 0 { buf = append(buf, tmp[:n]...) } if err != nil || len(buf) > 1024*1024 { break } if idx := indexOf(buf, '\n'); idx >= 0 { buf = buf[:idx] break } } conn.SetReadDeadline(time.Time{}) //nolint:errcheck var msg struct { Cmd string `json:"cmd"` Cwd string `json:"cwd"` Env map[string]string `json:"env"` Session string `json:"session"` } if err := json.Unmarshal(buf, &msg); err != nil { sendFrame(conn, FrameData, []byte(fmt.Sprintf("elevate: bad request: %v\n", err))) sendFrame(conn, FrameExit, []byte("1")) return } if msg.Cmd == "" { sendFrame(conn, FrameData, []byte("elevate: empty command\n")) sendFrame(conn, FrameExit, []byte("1")) return } // Identify caller session from request payload sessionID := msg.Session if sessionID == "" { b.logf("elevate: denied cmd=%q (no session identity)", msg.Cmd) sendFrame(conn, FrameData, []byte("elevate-broker: denied (no session identity)\n")) sendFrame(conn, FrameExit, []byte("1")) return } if b.sessionValid != nil && !b.sessionValid(sessionID) { b.logf("elevate: denied cmd=%q session=%s (session not found)", msg.Cmd, sessionID) sendFrame(conn, FrameData, []byte("elevate-broker: denied (session not found)\n")) sendFrame(conn, FrameExit, []byte("1")) return } // 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 (too many denied/unanswered requests)\n")) sendFrame(conn, FrameExit, []byte("1")) return } // Check policy (session policy includes global) effective := b.SessionPolicy(sessionID) if effective.Matches(msg.Cmd, "") { b.logf("elevate: auto-approved cmd=%q session=%s", msg.Cmd, sessionID) b.executeAndStream(conn, msg.Cmd, msg.Cwd, msg.Env) return } // Create pending request req := &Request{ ID: nextRequestID(), Cmd: msg.Cmd, Cwd: msg.Cwd, Env: msg.Env, SessionID: sessionID, CreatedAt: time.Now(), resolved: make(chan Resolution, 1), conn: conn, } b.mu.Lock() b.pending[req.ID] = req b.mu.Unlock() // Notify (D-Bus notification, etc.) b.notify(req) b.logf("elevate: pending id=%s cmd=%q session=%s", req.ID, msg.Cmd, sessionID) // Wait for resolution or timeout timer := time.NewTimer(RequestTTL) defer timer.Stop() var res Resolution select { case res = <-req.Resolved(): case <-timer.C: // Timeout - remove from pending and deny b.mu.Lock() delete(b.pending, req.ID) b.mu.Unlock() res = ResolveTimeout } switch res { case ResolveApprove, ResolvePersist: 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")) } } func (b *Broker) executeAndStream(conn net.Conn, cmd, cwd string, env map[string]string) { c := exec.Command("bash", "-c", cmd) c.Dir = cwd if len(env) > 0 { c.Env = make([]string, 0, len(env)) for k, v := range env { c.Env = append(c.Env, k+"="+v) } } c.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} // Pipe stdout+stderr to the connection via frames pr, pw, err := os.Pipe() if err != nil { sendFrame(conn, FrameData, []byte(fmt.Sprintf("elevate: pipe: %v\n", err))) sendFrame(conn, FrameExit, []byte("1")) return } c.Stdout = pw c.Stderr = pw if err := c.Start(); err != nil { pw.Close() pr.Close() sendFrame(conn, FrameData, []byte(fmt.Sprintf("elevate: exec: %v\n", err))) sendFrame(conn, FrameExit, []byte("1")) return } pw.Close() // Stream output buf := make([]byte, 4096) for { n, err := pr.Read(buf) if n > 0 { sendFrame(conn, FrameData, buf[:n]) } if err != nil { break } } pr.Close() err = c.Wait() exitCode := 0 if err != nil { if ee, ok := err.(*exec.ExitError); ok { exitCode = ee.ExitCode() } else { exitCode = 1 } } sendFrame(conn, FrameExit, []byte(fmt.Sprintf("%d", exitCode))) } func sendFrame(conn net.Conn, frameType byte, payload []byte) { header := make([]byte, 5) header[0] = frameType binary.BigEndian.PutUint32(header[1:], uint32(len(payload))) conn.Write(header) //nolint:errcheck conn.Write(payload) //nolint:errcheck } func indexOf(b []byte, c byte) int { for i, v := range b { if v == c { return i } } return -1 } func socketDir(path string) string { idx := strings.LastIndex(path, "/") if idx < 0 { return "." } return path[:idx] }