Files
rvbox/internal/agentproto/validate_test.go
T

265 lines
14 KiB
Go

package agentproto
import (
"bytes"
"crypto/sha256"
"errors"
"strings"
"testing"
"time"
"github.com/klauspost/compress/zstd"
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/timestamppb"
)
const testIssueUUID = "01890a5d-ac96-7a11-b234-123456789abc"
func TestDecodeEnvelopeBoundaries_HP_PROTO_01(t *testing.T) {
t.Parallel()
limits := DefaultLimits()
hello := &rvboxv1.AgentEnvelope{Payload: &rvboxv1.AgentEnvelope_ClientHello{ClientHello: validHello()}}
encoded, err := proto.Marshal(hello)
if err != nil {
t.Fatal(err)
}
decoded, err := DecodeEnvelope(encoded, limits, rvboxv1.Platform_PLATFORM_WINDOWS)
if err != nil || decoded.GetClientHello() == nil {
t.Fatalf("DecodeEnvelope(valid) = (%v, %v)", decoded, err)
}
if _, err := DecodeEnvelope(make([]byte, limits.MaxEnvelopeBytes+1), limits, rvboxv1.Platform_PLATFORM_WINDOWS); !errors.Is(err, ErrEnvelopeTooLarge) {
t.Fatalf("oversized error = %v", err)
}
if _, err := DecodeEnvelope([]byte{0xff}, limits, rvboxv1.Platform_PLATFORM_WINDOWS); !errors.Is(err, ErrInvalidEnvelope) {
t.Fatalf("malformed error = %v", err)
}
noPayload, _ := proto.Marshal(&rvboxv1.AgentEnvelope{})
if _, err := DecodeEnvelope(noPayload, limits, rvboxv1.Platform_PLATFORM_WINDOWS); !errors.Is(err, ErrInvalidEnvelope) {
t.Fatalf("missing payload error = %v", err)
}
}
func TestEnvelopeSessionFencingShape_BH_SES_01(t *testing.T) {
t.Parallel()
hello := &rvboxv1.AgentEnvelope{SessionId: "claimed", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ClientHello{ClientHello: validHello()}}
if err := ValidateEnvelope(hello, DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("ClientHello with session accepted")
}
welcome := &rvboxv1.AgentEnvelope{Payload: &rvboxv1.AgentEnvelope_ServerWelcome{ServerWelcome: &rvboxv1.ServerWelcome{SelectedProtocol: &rvboxv1.ProtocolVersion{Major: 1}, ServerTime: fixedTimestamp()}}}
if err := ValidateEnvelope(welcome, DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("non-hello without session accepted")
}
}
func TestExecutionSpecValidation_HP_PROTO_01(t *testing.T) {
t.Parallel()
limits := DefaultLimits()
valid := validCommandSpec()
if err := ValidateExecutionSpec(valid, limits, rvboxv1.Platform_PLATFORM_LINUX); err != nil {
t.Fatalf("valid spec: %v", err)
}
tests := []struct {
name string
mutate func(*rvboxv1.ExecutionSpec)
platform rvboxv1.Platform
}{
{"missing source", func(spec *rvboxv1.ExecutionSpec) { spec.Source = nil }, rvboxv1.Platform_PLATFORM_LINUX},
{"wrong shell", func(spec *rvboxv1.ExecutionSpec) { spec.ShellType = rvboxv1.ShellType_SHELL_CMD }, rvboxv1.Platform_PLATFORM_LINUX},
{"Unix elevation", func(spec *rvboxv1.ExecutionSpec) { spec.Elevated = true }, rvboxv1.Platform_PLATFORM_LINUX},
{"bad env key", func(spec *rvboxv1.ExecutionSpec) { spec.EnvOverrides = map[string]string{"A=B": "x"} }, rvboxv1.Platform_PLATFORM_LINUX},
{"duplicate profile", func(spec *rvboxv1.ExecutionSpec) {
spec.ExecutionProfiles = []rvboxv1.ExecutionProfile{rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_MEDIUM, rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_MEDIUM}
}, rvboxv1.Platform_PLATFORM_LINUX},
{"same dimension", func(spec *rvboxv1.ExecutionSpec) {
spec.ExecutionProfiles = []rvboxv1.ExecutionProfile{rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_MEDIUM, rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_HEAVY}
}, rvboxv1.Platform_PLATFORM_LINUX},
{"light combination", func(spec *rvboxv1.ExecutionSpec) {
spec.ExecutionProfiles = []rvboxv1.ExecutionProfile{rvboxv1.ExecutionProfile_EXECUTION_PROFILE_LIGHT, rvboxv1.ExecutionProfile_EXECUTION_PROFILE_MEM_HEAVY}
}, rvboxv1.Platform_PLATFORM_LINUX},
{"unspecified platform", func(spec *rvboxv1.ExecutionSpec) {}, rvboxv1.Platform_PLATFORM_UNSPECIFIED},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
t.Parallel()
spec := proto.Clone(valid).(*rvboxv1.ExecutionSpec)
test.mutate(spec)
if err := ValidateExecutionSpec(spec, limits, test.platform); err == nil {
t.Fatal("ValidateExecutionSpec() succeeded")
}
})
}
tooSmall := limits
tooSmall.MaxExecutionSpecBytes = uint64(proto.Size(valid) - 1)
if err := ValidateExecutionSpec(valid, tooSmall, rvboxv1.Platform_PLATFORM_LINUX); !errors.Is(err, ErrExecutionSpecTooLarge) {
t.Fatalf("spec size error = %v", err)
}
}
func TestScriptDescriptorHostileNames_BH_SCRIPT_01(t *testing.T) {
t.Parallel()
limits := DefaultLimits()
valid := &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_POWERSHELL, Source: &rvboxv1.ExecutionSpec_Script{Script: &rvboxv1.ScriptDescriptor{Filename: "update.ps1", SizeBytes: 1, Sha256: make([]byte, 32)}}}
if err := ValidateExecutionSpec(valid, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err != nil {
t.Fatalf("valid script: %v", err)
}
for _, name := range []string{"../x.ps1", `dir\x.ps1`, "CON.txt", "NUL", "bad\x00.ps1", ".", ".."} {
spec := proto.Clone(valid).(*rvboxv1.ExecutionSpec)
spec.GetScript().Filename = name
if err := ValidateExecutionSpec(spec, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Errorf("script filename %q accepted", name)
}
}
badDigest := proto.Clone(valid).(*rvboxv1.ExecutionSpec)
badDigest.GetScript().Sha256 = []byte{1}
if err := ValidateExecutionSpec(badDigest, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("wrong script digest length accepted")
}
trailing := proto.Clone(valid).(*rvboxv1.ExecutionSpec)
trailing.GetScript().Filename = "safe.ps1."
if err := ValidateExecutionSpec(trailing, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("unsafe Windows trailing dot accepted")
}
}
func TestWindowsEnvironmentCaseCollision_BH_LAUNCH_01(t *testing.T) {
t.Parallel()
spec := &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_CMD, EnvOverrides: map[string]string{"Path": "one", "PATH": "two"}, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo ok"}}
if err := ValidateExecutionSpec(spec, DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("case-colliding Windows environment accepted")
}
}
func TestCommandDispatchUUIDRevisionAndTimestamp_BH_IDEM_01(t *testing.T) {
t.Parallel()
dispatch := &rvboxv1.CommandDispatch{IssueUuid: testIssueUUID, CommandRevision: 1, TargetSessionGeneration: 2, IssueTime: fixedTimestamp(), Spec: validCommandSpec(), ImmutableRequestSha256: make([]byte, sha256.Size)}
envelope := &rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 2, Payload: &rvboxv1.AgentEnvelope_CommandDispatch{CommandDispatch: dispatch}}
if err := ValidateEnvelope(envelope, DefaultLimits(), rvboxv1.Platform_PLATFORM_LINUX); err != nil {
t.Fatalf("valid dispatch: %v", err)
}
for _, mutate := range []func(*rvboxv1.CommandDispatch){
func(value *rvboxv1.CommandDispatch) { value.IssueUuid = "not-a-uuid" },
func(value *rvboxv1.CommandDispatch) { value.CommandRevision = 0 },
func(value *rvboxv1.CommandDispatch) { value.TargetSessionGeneration = 0 },
func(value *rvboxv1.CommandDispatch) { value.ImmutableRequestSha256 = nil },
func(value *rvboxv1.CommandDispatch) {
value.IssueTime = timestamppb.New(time.Date(10000, 1, 1, 0, 0, 0, 0, time.UTC))
},
} {
copyEnvelope := proto.Clone(envelope).(*rvboxv1.AgentEnvelope)
mutate(copyEnvelope.GetCommandDispatch())
if err := ValidateEnvelope(copyEnvelope, DefaultLimits(), rvboxv1.Platform_PLATFORM_LINUX); err == nil {
t.Fatal("invalid dispatch accepted")
}
}
}
func TestOutputChunkBoundedDecompression_HP_PROTO_05(t *testing.T) {
t.Parallel()
const limit = 64 << 10
raw := bytes.Repeat([]byte("compressible-data"), 100)
encoder, err := zstd.NewWriter(nil, zstd.WithEncoderConcurrency(1))
if err != nil {
t.Fatal(err)
}
compressed := encoder.EncodeAll(raw, nil)
encoder.Close()
chunk := &rvboxv1.OutputChunk{Stream: rvboxv1.StreamKind_STREAM_STDOUT, Compression: rvboxv1.Compression_COMPRESSION_ZSTD, Data: compressed, UncompressedSize: uint64(len(raw)), CompressedSize: uint64(len(compressed))}
decoded, err := DecodeOutputChunk(chunk, limit)
if err != nil || !bytes.Equal(decoded, raw) {
t.Fatalf("DecodeOutputChunk() = (%d bytes, %v)", len(decoded), err)
}
none := &rvboxv1.OutputChunk{Stream: rvboxv1.StreamKind_STREAM_STDERR, Compression: rvboxv1.Compression_COMPRESSION_NONE, Data: []byte("x"), UncompressedSize: 1, CompressedSize: 1}
if decoded, err := DecodeOutputChunk(none, limit); err != nil || string(decoded) != "x" {
t.Fatalf("uncompressed chunk = (%q, %v)", decoded, err)
}
for _, invalid := range []*rvboxv1.OutputChunk{
nil,
{Stream: rvboxv1.StreamKind_STREAM_KIND_UNSPECIFIED, Compression: rvboxv1.Compression_COMPRESSION_NONE},
{Stream: rvboxv1.StreamKind_STREAM_STDOUT, Compression: rvboxv1.Compression_COMPRESSION_NONE, Data: []byte("x"), UncompressedSize: 2, CompressedSize: 1},
{Stream: rvboxv1.StreamKind_STREAM_STDOUT, Compression: rvboxv1.Compression_COMPRESSION_ZSTD, Data: []byte("bad"), UncompressedSize: 1, CompressedSize: 3},
} {
if _, err := DecodeOutputChunk(invalid, limit); !errors.Is(err, ErrInvalidOutputChunk) {
t.Errorf("invalid chunk error = %v", err)
}
}
bombRaw := bytes.Repeat([]byte("b"), limit+1)
encoder, _ = zstd.NewWriter(nil, zstd.WithEncoderConcurrency(1))
bomb := encoder.EncodeAll(bombRaw, nil)
encoder.Close()
bombChunk := &rvboxv1.OutputChunk{Stream: rvboxv1.StreamKind_STREAM_STDOUT, Compression: rvboxv1.Compression_COMPRESSION_ZSTD, Data: bomb, UncompressedSize: limit, CompressedSize: uint64(len(bomb))}
if _, err := DecodeOutputChunk(bombChunk, limit); !errors.Is(err, ErrInvalidOutputChunk) {
t.Fatalf("decompression bomb error = %v", err)
}
}
func TestCommandEventPayloadValidation_BH_OUTFLOW_01(t *testing.T) {
t.Parallel()
event := &rvboxv1.CommandEvent{IssueUuid: testIssueUUID, EventSeq: 1, ObservedAt: fixedTimestamp(), Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: &rvboxv1.LifecycleChange{Lifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, CommandRevision: 1}}}
envelope := &rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_CommandEvent{CommandEvent: event}}
if err := ValidateEnvelope(envelope, DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS); err != nil {
t.Fatalf("valid event: %v", err)
}
conflicting := proto.Clone(envelope).(*rvboxv1.AgentEnvelope)
conflicting.GetCommandEvent().GetLifecycle().Detail = strings.Repeat("x", int(DefaultLimits().MaxDetailBytes)+1)
if err := ValidateEnvelope(conflicting, DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("oversized lifecycle detail accepted")
}
}
func TestReconciliationRowsAndCapacity_BH_SES_01(t *testing.T) {
t.Parallel()
limits := DefaultLimits()
request := &rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ReconcileRequest{ReconcileRequest: &rvboxv1.ReconcileRequest{Targets: []*rvboxv1.ReconcileTarget{{IssueUuid: testIssueUUID, LastServerEventSeq: 1, CommandRevision: 1, ImmutableRequestSha256: make([]byte, 32)}}}}}
if err := ValidateEnvelope(request, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err != nil {
t.Fatalf("valid reconcile request: %v", err)
}
duplicate := proto.Clone(request).(*rvboxv1.AgentEnvelope)
duplicate.GetReconcileRequest().Targets = append(duplicate.GetReconcileRequest().Targets, proto.Clone(duplicate.GetReconcileRequest().Targets[0]).(*rvboxv1.ReconcileTarget))
if err := ValidateEnvelope(duplicate, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("duplicate reconcile target accepted")
}
snapshot := &rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ReconcileSnapshot{ReconcileSnapshot: &rvboxv1.ReconcileSnapshot{RetainedCommands: []*rvboxv1.ReconcileCommandState{{IssueUuid: testIssueUUID, Lifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, LastClientEventSeq: 1, CommandRevision: 1, ImmutableRequestSha256: make([]byte, 32)}}}}}
if err := ValidateEnvelope(snapshot, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err != nil {
t.Fatalf("valid snapshot: %v", err)
}
result := &rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ReconcileResult{ReconcileResult: &rvboxv1.ReconcileResult{TerminateLocalIssueUuids: []string{testIssueUUID}, DiscardLocalTerminalIssueUuids: []string{testIssueUUID}}}}
if err := ValidateEnvelope(result, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("same UUID in both reconcile result actions accepted")
}
capacity := &rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ClientCapacity{ClientCapacity: &rvboxv1.ClientCapacity{RunningCommands: 2, MaxRunningCommands: 1, MaxQueuedCommands: 1}}}
if err := ValidateEnvelope(capacity, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err == nil {
t.Fatal("capacity overflow accepted")
}
}
func validHello() *rvboxv1.ClientHello {
return &rvboxv1.ClientHello{ClientId: "host-01", SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, DaemonVersion: "test", Platform: rvboxv1.Platform_PLATFORM_WINDOWS, Architecture: "amd64", DaemonCwd: `C:\ProgramData\RVBox\work`, SupportedShells: []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_POWERSHELL}, ClientInstanceId: "instance-opaque", MaxRunningCommands: 16, MaxQueuedCommands: 100, SentAt: fixedTimestamp()}
}
func validCommandSpec() *rvboxv1.ExecutionSpec {
return &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_SH, EnvOverrides: map[string]string{"LANG": "C"}, ExecutionProfiles: []rvboxv1.ExecutionProfile{rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_MEDIUM, rvboxv1.ExecutionProfile_EXECUTION_PROFILE_MEM_MEDIUM}, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo ok"}}
}
func fixedTimestamp() *timestamppb.Timestamp {
return timestamppb.New(time.Unix(1_700_000_000, 0).UTC())
}