This repository has been archived on 2026-08-16. You can view files and clone it, but cannot push or open issues or pull requests.
ollie-core/pkg/elevate/broker.go

427 lines
10 KiB
Go

package elevate
import (
"encoding/binary"
"encoding/json"
"fmt"
"net"
"os"
"os/exec"
"strings"
"sync"
"syscall"
"time"
)
const (
RequestTTL = 300 * time.Second
FrameData = 'd'
FrameExit = 'x'
MaxPerTurn = 3 // max elevation requests per agent turn
)
// 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
turnMu sync.Mutex
turnCount map[string]int // sessionID -> requests this turn
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),
turnCount: make(map[string]int),
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.turnMu.Lock()
delete(b.turnCount, sessionID)
b.turnMu.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) {
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
}
// 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: max 3 elevation requests per agent turn
if !b.checkTurnLimit(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, 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.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]
}