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:
parent
c25f59ac7b
commit
2372c057bb
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Reference in New Issue