ollie/bypass/broker.go

534 lines
13 KiB
Go

package bypass
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)
)
// Broker manages elevation requests.
type Broker struct {
policy *PolicyStore
notify NotifyFunc
credential CredentialFunc
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
Credential CredentialFunc // prompts user for sudo password; nil = sudo unsupported
Logf func(string, ...any)
SessionValid func(sessionID string) bool // returns true if session exists
}
// NotifyFunc is called when a new elevation request needs user attention.
type NotifyFunc func(*Request)
// CredentialFunc prompts the user for a sudo password. Returns the password
// or an error if the user cancels. Called only after elevation is approved.
type CredentialFunc func(req *Request) (string, error)
// 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,
credential: cfg.Credential,
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"`
Sudo bool `json:"sudo"`
}
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)
if msg.Sudo {
b.executeWithSudo(conn, msg.Cmd, msg.Cwd, msg.Env, &Request{Cmd: msg.Cmd, SessionID: sessionID})
} else {
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,
Sudo: msg.Sudo,
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)
if req.Sudo {
b.executeWithSudo(conn, msg.Cmd, msg.Cwd, msg.Env, req)
} else {
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) executeWithSudo(conn net.Conn, cmd, cwd string, env map[string]string, req *Request) {
if b.credential == nil {
sendFrame(conn, FrameData, []byte("elevate-broker: sudo requested but no credential provider configured\n"))
sendFrame(conn, FrameExit, []byte("1"))
return
}
password, err := b.credential(req)
if err != nil {
b.logf("elevate: credential denied for cmd=%q: %v", cmd, err)
sendFrame(conn, FrameData, []byte(fmt.Sprintf("elevate-broker: credential denied: %v\n", err)))
sendFrame(conn, FrameExit, []byte("1"))
return
}
// Build: printf '%s\n' '<password>' | sudo -S bash -c '<cmd>'
// Password is first line on stdin; sudo -S reads it then runs the command.
c := exec.Command("sudo", "-S", "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 password to stdin
stdinPipe, err := c.StdinPipe()
if err != nil {
sendFrame(conn, FrameData, []byte(fmt.Sprintf("elevate: stdin pipe: %v\n", err)))
sendFrame(conn, FrameExit, []byte("1"))
return
}
// 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 sudo: %v\n", err)))
sendFrame(conn, FrameExit, []byte("1"))
return
}
// Send password to sudo -S
fmt.Fprintf(stdinPipe, "%s\n", password)
stdinPipe.Close()
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 (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]
}