From 2372c057bb7f82dd810ae9ebdf303e4621b9e47c Mon Sep 17 00:00:00 2001 From: Levi Neely Date: Wed, 29 Jul 2026 16:30:16 +0200 Subject: [PATCH] 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 --- pkg/elevate/broker.go | 82 ++++++++++++++++--------------------- pkg/elevate/broker_test.go | 53 ++++++++++++++++++++---- pkg/tools/execute/server.go | 9 ++-- 3 files changed, 87 insertions(+), 57 deletions(-) diff --git a/pkg/elevate/broker.go b/pkg/elevate/broker.go index e46be76..84b22d7 100644 --- a/pkg/elevate/broker.go +++ b/pkg/elevate/broker.go @@ -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 diff --git a/pkg/elevate/broker_test.go b/pkg/elevate/broker_test.go index b58b918..9d214f0 100644 --- a/pkg/elevate/broker_test.go +++ b/pkg/elevate/broker_test.go @@ -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") diff --git a/pkg/tools/execute/server.go b/pkg/tools/execute/server.go index 9332b52..2d88ea6 100644 --- a/pkg/tools/execute/server.go +++ b/pkg/tools/execute/server.go @@ -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()