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") } }