elevate: require valid session identity for all requests

- Client sends OLLIE_SESSION_ID in the request JSON payload
- Broker denies requests with no session identity
- Broker denies requests for sessions that don't exist
- Remove dead SO_PEERCRED logic (client is same process, useless)
- Remove allow-unknown bypass from checkTurnLimit
This commit is contained in:
Levi Neely 2026-07-29 16:30:16 +02:00
parent c25f59ac7b
commit 2372c057bb
3 changed files with 87 additions and 57 deletions

View File

@ -26,9 +26,10 @@ type NotifyFunc func(req *Request)
// Broker manages elevation requests.
type Broker struct {
policy *PolicyStore
notify NotifyFunc
logf func(string, ...any)
policy *PolicyStore
notify NotifyFunc
logf func(string, ...any)
sessionValid func(string) bool
mu sync.RWMutex
pending map[string]*Request // id -> request
@ -42,10 +43,11 @@ type Broker struct {
// BrokerConfig configures the broker.
type BrokerConfig struct {
SocketPath string
PolicyPath string // disk path for global policy YAML
Notify NotifyFunc
Logf func(string, ...any)
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.
@ -72,13 +74,14 @@ func NewBroker(cfg BrokerConfig) (*Broker, error) {
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),
turnCount: make(map[string]int),
listener: ln,
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()
@ -184,11 +187,8 @@ func (b *Broker) ResetTurn(sessionID string) {
// 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 {
if sessionID == "" {
// Unknown session — allow (can't enforce without identity)
return true
}
b.turnMu.Lock()
defer b.turnMu.Unlock()
n := b.turnCount[sessionID]
@ -251,9 +251,10 @@ func (b *Broker) handleConn(conn net.Conn) {
conn.SetReadDeadline(time.Time{}) //nolint:errcheck
var msg struct {
Cmd string `json:"cmd"`
Cwd string `json:"cwd"`
Env map[string]string `json:"env"`
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)))
@ -267,8 +268,20 @@ func (b *Broker) handleConn(conn net.Conn) {
return
}
// Identify caller session via SO_PEERCRED
sessionID := b.identifySession(conn)
// 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) {
@ -387,29 +400,6 @@ func (b *Broker) executeAndStream(conn net.Conn, cmd, cwd string, env map[string
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

View File

@ -7,6 +7,7 @@ import (
"net"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
@ -18,9 +19,10 @@ func startTestBroker(t *testing.T) (*Broker, string) {
policyPath := filepath.Join(dir, "policy.yaml")
b, err := NewBroker(BrokerConfig{
SocketPath: sockPath,
PolicyPath: policyPath,
Logf: t.Logf,
SocketPath: sockPath,
PolicyPath: policyPath,
Logf: t.Logf,
SessionValid: func(id string) bool { return id == "test-session" },
})
if err != nil {
t.Fatal(err)
@ -30,16 +32,21 @@ func startTestBroker(t *testing.T) (*Broker, string) {
}
func sendRequest(t *testing.T, sockPath, cmd, cwd string) net.Conn {
return sendRequestWithSession(t, sockPath, cmd, cwd, "test-session")
}
func sendRequestWithSession(t *testing.T, sockPath, cmd, cwd, session 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"}})
Cmd string `json:"cmd"`
Cwd string `json:"cwd"`
Env map[string]string `json:"env"`
Session string `json:"session,omitempty"`
}{Cmd: cmd, Cwd: cwd, Env: map[string]string{"PATH": "/usr/bin"}, Session: session})
req = append(req, '\n')
conn.Write(req) //nolint:errcheck
return conn
@ -91,6 +98,38 @@ func TestBrokerAutoApprove(t *testing.T) {
}
}
func TestBrokerDenyNoSession(t *testing.T) {
_, sockPath := startTestBroker(t)
// Send a request with no session ID
conn := sendRequestWithSession(t, sockPath, "echo hello", os.TempDir(), "")
defer conn.Close()
output, exitCode := readFrames(t, conn)
if exitCode != 1 {
t.Errorf("expected exit 1, got %d", exitCode)
}
if !strings.Contains(output, "no session identity") {
t.Errorf("expected 'no session identity' in output, got %q", output)
}
}
func TestBrokerDenyInvalidSession(t *testing.T) {
_, sockPath := startTestBroker(t)
// Send a request with a session ID that doesn't exist
conn := sendRequestWithSession(t, sockPath, "echo hello", os.TempDir(), "nonexistent-session")
defer conn.Close()
output, exitCode := readFrames(t, conn)
if exitCode != 1 {
t.Errorf("expected exit 1, got %d", exitCode)
}
if !strings.Contains(output, "session not found") {
t.Errorf("expected 'session not found' in output, got %q", output)
}
}
func TestBrokerDenyOnTimeout(t *testing.T) {
if testing.Short() {
t.Skip("skipping timeout test in short mode")

View File

@ -388,10 +388,11 @@ func (e *Server) executeElevated(ctx context.Context, cmd, dir string, timeout i
}
}
reqJSON, _ := json.Marshal(struct {
Cmd string `json:"cmd"`
Cwd string `json:"cwd"`
Env map[string]string `json:"env"`
}{Cmd: cmd, Cwd: dir, Env: envMap})
Cmd string `json:"cmd"`
Cwd string `json:"cwd"`
Env map[string]string `json:"env"`
Session string `json:"session,omitempty"`
}{Cmd: cmd, Cwd: dir, Env: envMap, Session: os.Getenv("OLLIE_SESSION_ID")})
reqJSON = append(reqJSON, '\n')
if _, err := conn.Write(reqJSON); err != nil {
conn.Close()