elevate: integrated elevation broker package

Core elevation broker that replaces the superpowerd adapter.
Handles command execution outside the sandbox with human-in-the-loop
approval via desktop notifications.

- Policy store: YAML-backed global policy + in-memory session policies
- Request lifecycle: 300s TTL, approve/deny/persist resolution
- Command execution: bash -c with caller env/cwd, streaming d/x frames
- SO_PEERCRED for caller identification
- NotifyFunc callback for UI integration (D-Bus, 9P, etc.)
This commit is contained in:
Levi Neely 2026-07-27 10:31:10 +02:00
parent d903980256
commit c418148dc8
6 changed files with 990 additions and 75 deletions

391
pkg/elevate/broker.go Normal file
View File

@ -0,0 +1,391 @@
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'
)
// 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)
mu sync.RWMutex
pending map[string]*Request // id -> request
sessions map[string]*Policy // sessionID -> session policy
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)
}
// 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,
pending: make(map[string]*Request),
sessions: make(map[string]*Policy),
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)
}
// 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"`
}
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 via SO_PEERCRED
sessionID := b.identifySession(conn)
// 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 (b *Broker) identifySession(conn net.Conn) string {
// Use SO_PEERCRED to get caller PID, then map to session.
// For now, return empty (can be wired up when session manager is available).
uc, ok := conn.(*net.UnixConn);
if !ok {
return ""
}
raw, err := uc.SyscallConn()
if err != nil {
return ""
}
var cred *syscall.Ucred
raw.Control(func(fd uintptr) { //nolint:errcheck
cred, _ = syscall.GetsockoptUcred(int(fd), syscall.SOL_SOCKET, syscall.SO_PEERCRED)
})
if cred == nil {
return ""
}
// TODO: map PID to session via session manager
_ = cred.Pid
return ""
}
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]
}

198
pkg/elevate/broker_test.go Normal file
View File

@ -0,0 +1,198 @@
package elevate
import (
"encoding/binary"
"encoding/json"
"fmt"
"net"
"os"
"path/filepath"
"testing"
"time"
)
func startTestBroker(t *testing.T) (*Broker, string) {
t.Helper()
dir := t.TempDir()
sockPath := filepath.Join(dir, "elevate.sock")
policyPath := filepath.Join(dir, "policy.yaml")
b, err := NewBroker(BrokerConfig{
SocketPath: sockPath,
PolicyPath: policyPath,
Logf: t.Logf,
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(b.Close)
return b, sockPath
}
func sendRequest(t *testing.T, sockPath, cmd, cwd string) net.Conn {
t.Helper()
conn, err := net.DialTimeout("unix", sockPath, 2*time.Second)
if err != nil {
t.Fatal(err)
}
req, _ := json.Marshal(struct {
Cmd string `json:"cmd"`
Cwd string `json:"cwd"`
Env map[string]string `json:"env"`
}{Cmd: cmd, Cwd: cwd, Env: map[string]string{"PATH": "/usr/bin"}})
req = append(req, '\n')
conn.Write(req) //nolint:errcheck
return conn
}
func readFrames(t *testing.T, conn net.Conn) (string, int) {
t.Helper()
var output string
exitCode := -1
header := make([]byte, 5)
for {
conn.SetReadDeadline(time.Now().Add(5 * time.Second)) //nolint:errcheck
_, err := conn.Read(header)
if err != nil {
break
}
length := binary.BigEndian.Uint32(header[1:5])
payload := make([]byte, length)
if length > 0 {
conn.Read(payload) //nolint:errcheck
}
switch header[0] {
case 'd':
output += string(payload)
case 'x':
fmt.Sscanf(string(payload), "%d", &exitCode)
return output, exitCode
}
}
return output, exitCode
}
func TestBrokerAutoApprove(t *testing.T) {
b, sockPath := startTestBroker(t)
// Add a policy rule
b.GlobalPolicy().AddGlobal(Rule{Cmd: "echo *"})
// Send a matching command
conn := sendRequest(t, sockPath, "echo hello", os.TempDir())
defer conn.Close()
output, exitCode := readFrames(t, conn)
if exitCode != 0 {
t.Errorf("expected exit 0, got %d; output: %s", exitCode, output)
}
if output != "hello\n" {
t.Errorf("expected 'hello\n', got %q", output)
}
}
func TestBrokerDenyOnTimeout(t *testing.T) {
if testing.Short() {
t.Skip("skipping timeout test in short mode")
}
// Use a very short TTL for testing
dir := t.TempDir()
sockPath := filepath.Join(dir, "elevate.sock")
policyPath := filepath.Join(dir, "policy.yaml")
b, err := NewBroker(BrokerConfig{
SocketPath: sockPath,
PolicyPath: policyPath,
Logf: t.Logf,
})
if err != nil {
t.Fatal(err)
}
defer b.Close()
// No policy match → goes to pending
conn := sendRequest(t, sockPath, "rm -rf /", os.TempDir())
defer conn.Close()
// Manually deny after a short wait
time.Sleep(100 * time.Millisecond)
pending := b.Pending()
if len(pending) != 1 {
t.Fatalf("expected 1 pending, got %d", len(pending))
}
b.Resolve(pending[0].ID, ResolveDeny)
output, exitCode := readFrames(t, conn)
if exitCode != 1 {
t.Errorf("expected exit 1 on deny, got %d", exitCode)
}
_ = output
}
func TestBrokerApproveAndExecute(t *testing.T) {
b, sockPath := startTestBroker(t)
conn := sendRequest(t, sockPath, "echo approved", os.TempDir())
defer conn.Close()
// Wait for it to appear in pending
time.Sleep(100 * time.Millisecond)
pending := b.Pending()
if len(pending) != 1 {
t.Fatalf("expected 1 pending, got %d", len(pending))
}
// Approve it
b.Resolve(pending[0].ID, ResolveApprove)
output, exitCode := readFrames(t, conn)
if exitCode != 0 {
t.Errorf("expected exit 0, got %d; output: %s", exitCode, output)
}
if output != "approved\n" {
t.Errorf("expected 'approved\n', got %q", output)
}
}
func TestBrokerPersist(t *testing.T) {
b, sockPath := startTestBroker(t)
// First request: no policy, goes to pending
conn1 := sendRequest(t, sockPath, "echo persist-me", os.TempDir())
defer conn1.Close()
time.Sleep(100 * time.Millisecond)
pending := b.Pending()
if len(pending) != 1 {
t.Fatalf("expected 1 pending, got %d", len(pending))
}
// Persist it
b.Resolve(pending[0].ID, ResolvePersist)
output, exitCode := readFrames(t, conn1)
if exitCode != 0 {
t.Errorf("first request: expected exit 0, got %d", exitCode)
}
if output != "persist-me\n" {
t.Errorf("first request: expected 'persist-me\n', got %q", output)
}
// Second request with same command: should auto-approve (no pending)
conn2 := sendRequest(t, sockPath, "echo persist-me", os.TempDir())
defer conn2.Close()
output2, exitCode2 := readFrames(t, conn2)
if exitCode2 != 0 {
t.Errorf("second request: expected exit 0, got %d", exitCode2)
}
if output2 != "persist-me\n" {
t.Errorf("second request: expected 'persist-me\n', got %q", output2)
}
// Verify no pending requests
if len(b.Pending()) != 0 {
t.Error("expected no pending after persist")
}
}

125
pkg/elevate/policy.go Normal file
View File

@ -0,0 +1,125 @@
// Package elevate implements the elevation broker for running commands
// outside the Landlock sandbox with human-in-the-loop approval.
package elevate
import (
"os"
"path/filepath"
"sync"
"gopkg.in/yaml.v3"
)
// Rule represents a single auto-approve rule.
type Rule struct {
Cmd string `yaml:"cmd,omitempty"` // command pattern (glob)
SSH string `yaml:"ssh,omitempty"` // SSH key fingerprint
}
// Match returns true if the rule matches the given command or fingerprint.
func (r Rule) Match(cmd, fingerprint string) bool {
if r.Cmd != "" {
matched, _ := filepath.Match(r.Cmd, cmd)
if matched {
return true
}
// Also try exact match (glob patterns may not cover all cases)
if r.Cmd == cmd {
return true
}
}
if r.SSH != "" && r.SSH == fingerprint {
return true
}
return false
}
// Policy is a set of auto-approve rules.
type Policy struct {
Rules []Rule `yaml:"rules"`
}
// Matches returns true if any rule matches.
func (p *Policy) Matches(cmd, fingerprint string) bool {
for _, r := range p.Rules {
if r.Match(cmd, fingerprint) {
return true
}
}
return false
}
// Add appends a rule and returns true if it was new.
func (p *Policy) Add(r Rule) bool {
for _, existing := range p.Rules {
if existing == r {
return false
}
}
p.Rules = append(p.Rules, r)
return true
}
// Marshal returns the YAML representation.
func (p *Policy) Marshal() ([]byte, error) {
return yaml.Marshal(p)
}
// PolicyStore manages a disk-backed global policy and in-memory session policies.
type PolicyStore struct {
mu sync.RWMutex
global Policy
path string // disk path for global policy
}
// NewPolicyStore loads (or creates) the global policy from disk.
func NewPolicyStore(path string) *PolicyStore {
ps := &PolicyStore{path: path}
data, err := os.ReadFile(path)
if err == nil {
yaml.Unmarshal(data, &ps.global) //nolint:errcheck
}
return ps
}
// Global returns a copy of the global policy.
func (ps *PolicyStore) Global() Policy {
ps.mu.RLock()
defer ps.mu.RUnlock()
cp := Policy{Rules: make([]Rule, len(ps.global.Rules))}
copy(cp.Rules, ps.global.Rules)
return cp
}
// SetGlobal replaces the global policy and saves to disk.
func (ps *PolicyStore) SetGlobal(p Policy) error {
ps.mu.Lock()
defer ps.mu.Unlock()
ps.global = p
return ps.save()
}
// AddGlobal adds a rule to the global policy and saves.
func (ps *PolicyStore) AddGlobal(r Rule) error {
ps.mu.Lock()
defer ps.mu.Unlock()
ps.global.Add(r)
return ps.save()
}
// MatchesGlobal checks if the global policy matches.
func (ps *PolicyStore) MatchesGlobal(cmd, fingerprint string) bool {
ps.mu.RLock()
defer ps.mu.RUnlock()
return ps.global.Matches(cmd, fingerprint)
}
func (ps *PolicyStore) save() error {
data, err := yaml.Marshal(&ps.global)
if err != nil {
return err
}
dir := filepath.Dir(ps.path)
os.MkdirAll(dir, 0700) //nolint:errcheck
return os.WriteFile(ps.path, data, 0600)
}

104
pkg/elevate/policy_test.go Normal file
View File

@ -0,0 +1,104 @@
package elevate
import (
"os"
"path/filepath"
"testing"
)
func TestRuleMatch(t *testing.T) {
tests := []struct {
rule Rule
cmd string
fingerprint string
want bool
}{
{Rule{Cmd: "echo hello"}, "echo hello", "", true},
{Rule{Cmd: "echo hello"}, "echo world", "", false},
{Rule{Cmd: "git *"}, "git push", "", true},
{Rule{Cmd: "git *"}, "git pull", "", true},
{Rule{Cmd: "git *"}, "make build", "", false},
{Rule{Cmd: "echo *"}, "echo hello world", "", true}, // filepath.Match * matches any non-separator chars
{Rule{SSH: "SHA256:abc123"}, "", "SHA256:abc123", true},
{Rule{SSH: "SHA256:abc123"}, "", "SHA256:other", false},
{Rule{}, "anything", "anything", false},
}
for _, tt := range tests {
got := tt.rule.Match(tt.cmd, tt.fingerprint)
if got != tt.want {
t.Errorf("Rule%+v.Match(%q, %q) = %v; want %v", tt.rule, tt.cmd, tt.fingerprint, got, tt.want)
}
}
}
func TestPolicyMatches(t *testing.T) {
p := &Policy{
Rules: []Rule{
{Cmd: "echo *"},
{Cmd: "git push"},
{SSH: "SHA256:mykey"},
},
}
if !p.Matches("git push", "") {
t.Error("expected git push to match")
}
if !p.Matches("", "SHA256:mykey") {
t.Error("expected SSH key to match")
}
if p.Matches("rm -rf /", "") {
t.Error("expected rm to not match")
}
}
func TestPolicyAdd(t *testing.T) {
p := &Policy{}
if !p.Add(Rule{Cmd: "echo hi"}) {
t.Error("first add should return true")
}
if p.Add(Rule{Cmd: "echo hi"}) {
t.Error("duplicate add should return false")
}
if len(p.Rules) != 1 {
t.Errorf("expected 1 rule, got %d", len(p.Rules))
}
}
func TestPolicyStorePersistence(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "policy.yaml")
// Create and save
ps := NewPolicyStore(path)
ps.AddGlobal(Rule{Cmd: "make build"})
ps.AddGlobal(Rule{SSH: "SHA256:testkey"})
if !ps.MatchesGlobal("make build", "") {
t.Error("expected make build to match")
}
// Reload from disk
ps2 := NewPolicyStore(path)
if !ps2.MatchesGlobal("make build", "") {
t.Error("expected make build to match after reload")
}
if !ps2.MatchesGlobal("", "SHA256:testkey") {
t.Error("expected SSH key to match after reload")
}
// Verify file exists
if _, err := os.Stat(path); err != nil {
t.Errorf("policy file not created: %v", err)
}
}
func TestPolicyStoreEmpty(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "nonexistent.yaml")
ps := NewPolicyStore(path)
if ps.MatchesGlobal("anything", "") {
t.Error("empty policy should not match")
}
}

69
pkg/elevate/request.go Normal file
View File

@ -0,0 +1,69 @@
package elevate
import (
"fmt"
"net"
"sync/atomic"
"time"
)
// Resolution is the outcome of a pending request.
type Resolution int
const (
ResolvePending Resolution = iota
ResolveApprove
ResolveDeny
ResolvePersist // approve + add to session policy
ResolveTimeout
)
func (r Resolution) String() string {
switch r {
case ResolveApprove:
return "approve"
case ResolveDeny:
return "deny"
case ResolvePersist:
return "persist"
case ResolveTimeout:
return "timeout"
default:
return "pending"
}
}
// Request represents a pending elevation request.
type Request struct {
ID string
Cmd string
Cwd string
Env map[string]string
SessionID string // ollie session that made the request (from SO_PEERCRED mapping)
CreatedAt time.Time
// Resolution channel — exactly one value sent when resolved.
resolved chan Resolution
conn net.Conn // the elevate client connection (held open until resolved)
}
// Resolved returns the channel that receives the resolution.
func (r *Request) Resolved() <-chan Resolution {
return r.resolved
}
// Summary returns a human-readable description for notifications.
func (r *Request) Summary() string {
sess := r.SessionID
if sess == "" {
sess = "unknown"
}
return fmt.Sprintf("[%s] %s\ncwd: %s", sess, r.Cmd, r.Cwd)
}
var requestCounter atomic.Uint64
func nextRequestID() string {
n := requestCounter.Add(1)
return fmt.Sprintf("%d", n)
}

View File

@ -3,9 +3,11 @@ package execute
import (
"bytes"
"context"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"os"
"os/exec"
"path/filepath"
@ -405,13 +407,17 @@ func (e *Server) Close() {
}
// executeElevated runs cmd outside the sandbox via x/elevate.
// Returns (output, error); a non-zero exit code is treated as an error.
// If detach is true, the process is immediately backgrounded into the detach registry.
// executeElevated runs cmd outside the sandbox via the integrated elevation broker.
// Connects to the broker socket, sends the request with the current env,
// and streams the framed response back.
func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout int, detach ...bool) (string, error) {
script := filepath.Join(PluginsPath(), "elevate")
if _, err := os.Stat(script); err != nil {
return "", fmt.Errorf("elevation not available: x/elevate not found")
sockPath := os.Getenv("OLLIE_ELEVATE_SOCKET")
if sockPath == "" {
xdg := os.Getenv("XDG_RUNTIME_DIR")
if xdg == "" {
return "", fmt.Errorf("elevation not available: no XDG_RUNTIME_DIR")
}
sockPath = filepath.Join(xdg, "ollie", "elevate.sock")
}
wantDetach := len(detach) > 0 && detach[0]
@ -423,25 +429,31 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
ctx, cancel = context.WithTimeout(ctx, time.Duration(timeout)*time.Second)
}
c := exec.CommandContext(ctx, script, "--", cmd)
c.Dir = dir
c.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
c.Cancel = func() error {
if c.Process != nil {
syscall.Kill(-c.Process.Pid, syscall.SIGKILL)
}
return nil
// Connect to broker
conn, err := net.DialTimeout("unix", sockPath, 5*time.Second)
if err != nil {
cancel()
return "", fmt.Errorf("elevation not available: %w", err)
}
c.WaitDelay = time.Second
var outputBuf bytes.Buffer
lw := &limitedWriter{
w: &outputBuf,
limit: 10 * 1024 * 1024,
stream: tools.StreamFunc(ctx),
// Send request
envMap := make(map[string]string, len(os.Environ()))
for _, kv := range os.Environ() {
if k, v, ok := strings.Cut(kv, "="); ok {
envMap[k] = v
}
}
reqJSON, _ := json.Marshal(struct {
Cmd string `json:"cmd"`
Cwd string `json:"cwd"`
Env map[string]string `json:"env"`
}{Cmd: cmd, Cwd: dir, Env: envMap})
reqJSON = append(reqJSON, '\n')
if _, err := conn.Write(reqJSON); err != nil {
conn.Close()
cancel()
return "", fmt.Errorf("elevated execution failed: write: %w", err)
}
c.Stdout = lw
c.Stderr = lw
// Set up detach channel
detachCh := make(chan struct{}, 1)
@ -454,67 +466,71 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
e.detachMu.Unlock()
}()
if err := c.Start(); err != nil {
cancel()
return "", fmt.Errorf("elevated execution failed: %v", err)
}
if wantDetach {
close(detachCh)
}
waitCh := make(chan error, 1)
go func() { waitCh <- c.Wait() }()
select {
case err := <-waitCh:
cancel()
combined := outputBuf.String()
if lw.truncated {
combined += "\n[output truncated at 10MB]"
}
if err != nil {
var exitErr *exec.ExitError
if errors.As(err, &exitErr) {
if combined == "" {
return "", fmt.Errorf("elevated execution failed (exit %d) with no output", exitErr.ExitCode())
// readFrames reads from the connection, writing to w and streaming.
// Returns exit code when 'x' frame arrives or -1 on error.
readFrames := func(w io.Writer, stream func(string)) int {
header := make([]byte, 5)
for {
conn.SetReadDeadline(time.Now().Add(1 * time.Second)) //nolint:errcheck
_, err := io.ReadFull(conn, header)
if err != nil {
if ne, ok := err.(net.Error); ok && ne.Timeout() {
// Check for context cancellation
select {
case <-ctx.Done():
return -1
default:
continue
}
}
return combined, fmt.Errorf("elevated execution failed (exit %d)\nOutput: %s", exitErr.ExitCode(), combined)
return -1
}
return combined, fmt.Errorf("elevated execution failed: %v", err)
}
return combined, nil
case <-ctx.Done():
cancel()
syscall.Kill(-c.Process.Pid, syscall.SIGKILL)
<-waitCh
combined := outputBuf.String()
if ctx.Err() == context.DeadlineExceeded {
return combined, fmt.Errorf("execution timeout after %d seconds", timeout)
}
return combined, fmt.Errorf("execution interrupted")
frameType := header[0]
length := binary.BigEndian.Uint32(header[1:5])
payload := make([]byte, length)
if length > 0 {
conn.SetReadDeadline(time.Now().Add(30 * time.Second)) //nolint:errcheck
if _, err := io.ReadFull(conn, payload); err != nil {
return -1
}
}
switch frameType {
case 'd':
w.Write(payload) //nolint:errcheck
if stream != nil {
stream(string(payload))
}
case 'x':
var code int
fmt.Sscanf(string(payload), "%d", &code)
return code
}
}
}
// Check if detach was requested
select {
case <-detachCh:
c.Cancel = nil
c.WaitDelay = 0
cancel()
// Detach: background the socket read into a goroutine
cmdStr := "elevated: " + cmd
if len(cmdStr) > 80 {
cmdStr = cmdStr[:77] + "..."
}
ring := newRingBuffer(ringBufSize)
lw.mu.Lock()
lw.w = ring
lw.mu.Unlock()
pid := int(time.Now().UnixNano() & 0x7FFFFFFF) // synthetic PID
proc := &DetachedProcess{
PID: c.Process.Pid,
PID: pid,
Command: cmdStr,
Started: time.Now(),
ring: ring,
cmd: c.Process,
done: make(chan struct{}),
}
e.detachMu.Lock()
@ -522,16 +538,12 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
e.detachMu.Unlock()
go func() {
waitErr := <-waitCh
defer conn.Close()
defer cancel()
exitCode := readFrames(ring, nil)
proc.mu.Lock()
proc.Exited = true
if waitErr != nil {
if exitErr, ok := waitErr.(*exec.ExitError); ok {
proc.ExitCode = exitErr.ExitCode()
} else {
proc.ExitCode = -1
}
}
proc.ExitCode = exitCode
proc.mu.Unlock()
close(proc.done)
if e.OnExit != nil {
@ -542,9 +554,25 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
if e.OnDetach != nil {
e.OnDetach(proc.PID, cmdStr)
}
return fmt.Sprintf("[detached: pid %d]", pid), nil
partial := outputBuf.String()
return partial + fmt.Sprintf("\n[detached: pid %d]", proc.PID), nil
default:
// Normal (foreground) execution
defer conn.Close()
defer cancel()
var outputBuf bytes.Buffer
streamFn := tools.StreamFunc(ctx)
exitCode := readFrames(&outputBuf, streamFn)
combined := outputBuf.String()
if exitCode != 0 {
if combined == "" {
return "", fmt.Errorf("elevated execution failed (exit %d)", exitCode)
}
return combined, fmt.Errorf("elevated execution failed (exit %d)\nOutput: %s", exitCode, combined)
}
return combined, nil
}
}