435 lines
11 KiB
Go
435 lines
11 KiB
Go
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]
|
|
}
|