309 lines
7.5 KiB
Go
309 lines
7.5 KiB
Go
package elevate
|
|
|
|
import (
|
|
"encoding/binary"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"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,
|
|
SessionValid: func(id string) bool { return id == "test-session" },
|
|
})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
t.Cleanup(b.Close)
|
|
return b, sockPath
|
|
}
|
|
|
|
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"`
|
|
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
|
|
}
|
|
|
|
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 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")
|
|
}
|
|
|
|
// 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")
|
|
}
|
|
}
|
|
|
|
func TestBrokerRateLimitLeakyBucket(t *testing.T) {
|
|
b, _ := startTestBroker(t)
|
|
|
|
sessionID := "test-session"
|
|
|
|
// No denials — should be allowed
|
|
for i := 0; i < BurstLimit+1; i++ {
|
|
if !b.checkRateLimit(sessionID) {
|
|
t.Fatalf("request %d should not be rate-limited (no denials yet)", i+1)
|
|
}
|
|
}
|
|
|
|
// Record BurstLimit denials — exhaust the bucket
|
|
for i := 0; i < BurstLimit; i++ {
|
|
b.recordDenial(sessionID)
|
|
}
|
|
|
|
// Now should be rate-limited
|
|
if b.checkRateLimit(sessionID) {
|
|
t.Error("should be rate-limited after BurstLimit denials")
|
|
}
|
|
}
|
|
|
|
func TestBrokerRateLimitRefill(t *testing.T) {
|
|
b, _ := startTestBroker(t)
|
|
|
|
sid := "refill-test"
|
|
|
|
// Exhaust the bucket
|
|
for i := 0; i < BurstLimit; i++ {
|
|
b.recordDenial(sid)
|
|
}
|
|
if b.checkRateLimit(sid) {
|
|
t.Fatal("should be rate-limited")
|
|
}
|
|
|
|
// Manually set the limiter to allow (simulates time passing)
|
|
// We can't easily wait 10s in a unit test, so verify RemoveSession resets.
|
|
b.RemoveSession(sid)
|
|
if !b.checkRateLimit(sid) {
|
|
t.Fatal("should be allowed after RemoveSession (fresh limiter)")
|
|
}
|
|
}
|
|
|
|
func TestBrokerRemoveSessionCleansUpLimiter(t *testing.T) {
|
|
b, _ := startTestBroker(t)
|
|
|
|
sid := "cleanup-test"
|
|
|
|
// Record some denials
|
|
b.recordDenial(sid)
|
|
b.recordDenial(sid)
|
|
|
|
// Remove session
|
|
b.RemoveSession(sid)
|
|
|
|
// Fresh limiter — exhaust again
|
|
for i := 0; i < BurstLimit; i++ {
|
|
b.recordDenial(sid)
|
|
}
|
|
if b.checkRateLimit(sid) {
|
|
t.Fatal("should be rate-limited after BurstLimit denials")
|
|
}
|
|
|
|
// Clean up and verify fresh state
|
|
b.RemoveSession(sid)
|
|
if !b.checkRateLimit(sid) {
|
|
t.Fatal("should be allowed after RemoveSession")
|
|
}
|
|
}
|