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:
parent
d903980256
commit
c418148dc8
|
|
@ -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]
|
||||
}
|
||||
|
|
@ -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")
|
||||
}
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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")
|
||||
}
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
Reference in New Issue