// Package agentproto validates generated Agent protocol DTOs at the wire edge. package agentproto import ( "crypto/sha256" "errors" "fmt" "path/filepath" "strings" "unicode/utf8" rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/domain" "google.golang.org/protobuf/proto" ) var ( ErrEnvelopeTooLarge = errors.New("decoded agent envelope exceeds limit") ErrExecutionSpecTooLarge = errors.New("serialized execution spec exceeds limit") ErrInvalidEnvelope = errors.New("invalid agent envelope") ErrInvalidExecutionSpec = errors.New("invalid execution spec") ) type Limits struct { MaxEnvelopeBytes uint64 MaxExecutionSpecBytes uint64 MaxRawChunkBytes uint64 MaxScriptBytes uint64 MaxDetailBytes uint64 } func DefaultLimits() Limits { return Limits{MaxEnvelopeBytes: 1 << 20, MaxExecutionSpecBytes: 768 << 10, MaxRawChunkBytes: 64 << 10, MaxScriptBytes: 10 << 20, MaxDetailBytes: 4 << 10} } // DecodeEnvelope checks the binary message size before protobuf allocation. func DecodeEnvelope(data []byte, limits Limits, peerPlatform rvboxv1.Platform) (*rvboxv1.AgentEnvelope, error) { if uint64(len(data)) > limits.MaxEnvelopeBytes { return nil, ErrEnvelopeTooLarge } var envelope rvboxv1.AgentEnvelope if err := (proto.UnmarshalOptions{DiscardUnknown: false}).Unmarshal(data, &envelope); err != nil { return nil, fmt.Errorf("%w: decode protobuf: %v", ErrInvalidEnvelope, err) } if err := ValidateEnvelope(&envelope, limits, peerPlatform); err != nil { return nil, err } return &envelope, nil } func ValidateEnvelope(envelope *rvboxv1.AgentEnvelope, limits Limits, peerPlatform rvboxv1.Platform) error { if envelope == nil || envelope.Payload == nil { return fmt.Errorf("%w: missing recognized payload", ErrInvalidEnvelope) } if uint64(proto.Size(envelope)) > limits.MaxEnvelopeBytes { return ErrEnvelopeTooLarge } _, hello := envelope.Payload.(*rvboxv1.AgentEnvelope_ClientHello) if hello { if envelope.SessionId != "" || envelope.SessionGeneration != 0 { return fmt.Errorf("%w: ClientHello must not claim a session", ErrInvalidEnvelope) } } else if envelope.SessionId == "" || envelope.SessionGeneration == 0 || len(envelope.SessionId) > 128 || !utf8.ValidString(envelope.SessionId) || strings.IndexByte(envelope.SessionId, 0) >= 0 { return fmt.Errorf("%w: session ID/generation required", ErrInvalidEnvelope) } switch payload := envelope.Payload.(type) { case *rvboxv1.AgentEnvelope_ClientHello: return validateClientHello(payload.ClientHello) case *rvboxv1.AgentEnvelope_ServerWelcome: if payload.ServerWelcome == nil || payload.ServerWelcome.SelectedProtocol == nil || payload.ServerWelcome.ServerTime == nil || payload.ServerWelcome.ServerTime.CheckValid() != nil { return fmt.Errorf("%w: invalid ServerWelcome", ErrInvalidEnvelope) } return validateProtocolVersion(payload.ServerWelcome.SelectedProtocol) case *rvboxv1.AgentEnvelope_CommandDispatch: return validateCommandDispatch(payload.CommandDispatch, limits, peerPlatform) case *rvboxv1.AgentEnvelope_CommandAccepted: return validateCommandAccepted(payload.CommandAccepted, limits) case *rvboxv1.AgentEnvelope_CommandEvent: return validateCommandEvent(payload.CommandEvent, limits) case *rvboxv1.AgentEnvelope_EventAck: if payload.EventAck == nil { return fmt.Errorf("%w: missing event acknowledgement", ErrInvalidEnvelope) } return validateIssueAndSequence(payload.EventAck.GetIssueUuid(), payload.EventAck.GetThroughEventSeq()) case *rvboxv1.AgentEnvelope_StdinWrite: if payload.StdinWrite == nil || payload.StdinWrite.WriteSeq == 0 || uint64(len(payload.StdinWrite.Data)) > limits.MaxRawChunkBytes { return fmt.Errorf("%w: invalid stdin write", ErrInvalidEnvelope) } return validateIssueUUID(payload.StdinWrite.IssueUuid) case *rvboxv1.AgentEnvelope_CloseStdin: if payload.CloseStdin == nil || payload.CloseStdin.WriteSeq == 0 { return fmt.Errorf("%w: invalid close stdin", ErrInvalidEnvelope) } return validateIssueUUID(payload.CloseStdin.IssueUuid) case *rvboxv1.AgentEnvelope_SignalCommand: if payload.SignalCommand == nil || payload.SignalCommand.CommandRevision == 0 || !validSignal(payload.SignalCommand.Signal) { return fmt.Errorf("%w: invalid signal request", ErrInvalidEnvelope) } return validateIssueUUID(payload.SignalCommand.IssueUuid) case *rvboxv1.AgentEnvelope_ScriptChunk: if payload.ScriptChunk == nil || len(payload.ScriptChunk.Sha256) != sha256.Size || uint64(len(payload.ScriptChunk.Data)) > limits.MaxRawChunkBytes || payload.ScriptChunk.Offset > limits.MaxScriptBytes || uint64(len(payload.ScriptChunk.Data)) > limits.MaxScriptBytes-payload.ScriptChunk.Offset { return fmt.Errorf("%w: invalid script chunk", ErrInvalidEnvelope) } return validateIssueUUID(payload.ScriptChunk.IssueUuid) case *rvboxv1.AgentEnvelope_ScriptCommit: if payload.ScriptCommit == nil || payload.ScriptCommit.SizeBytes > limits.MaxScriptBytes || len(payload.ScriptCommit.Sha256) != sha256.Size { return fmt.Errorf("%w: invalid script commit", ErrInvalidEnvelope) } return validateIssueUUID(payload.ScriptCommit.IssueUuid) case *rvboxv1.AgentEnvelope_ReconcileRequest: return validateReconcileRequest(payload.ReconcileRequest) case *rvboxv1.AgentEnvelope_ReconcileSnapshot: return validateReconcileSnapshot(payload.ReconcileSnapshot) case *rvboxv1.AgentEnvelope_ReconcileResult: return validateReconcileResult(payload.ReconcileResult) case *rvboxv1.AgentEnvelope_ClientCapacity: if payload.ClientCapacity == nil || payload.ClientCapacity.MaxRunningCommands == 0 || payload.ClientCapacity.MaxQueuedCommands == 0 || payload.ClientCapacity.RunningCommands > payload.ClientCapacity.MaxRunningCommands || payload.ClientCapacity.QueuedCommands > payload.ClientCapacity.MaxQueuedCommands { return fmt.Errorf("%w: invalid client capacity", ErrInvalidEnvelope) } return nil case *rvboxv1.AgentEnvelope_Error: if payload.Error == nil || payload.Error.Error == nil || payload.Error.Error.Code == rvboxv1.ControlError_CODE_UNSPECIFIED || uint64(len(payload.Error.Error.Message)) > limits.MaxDetailBytes { return fmt.Errorf("%w: invalid agent error", ErrInvalidEnvelope) } return nil default: return fmt.Errorf("%w: unsupported payload", ErrInvalidEnvelope) } } func ValidateExecutionSpec(spec *rvboxv1.ExecutionSpec, limits Limits, peerPlatform rvboxv1.Platform) error { if spec == nil || spec.Source == nil { return fmt.Errorf("%w: exactly one command or script source is required", ErrInvalidExecutionSpec) } if uint64(proto.Size(spec)) > limits.MaxExecutionSpecBytes { return ErrExecutionSpecTooLarge } if !validShell(spec.ShellType, peerPlatform) { return fmt.Errorf("%w: unsupported shell %s", ErrInvalidExecutionSpec, spec.ShellType) } if spec.Elevated && peerPlatform != rvboxv1.Platform_PLATFORM_WINDOWS { return fmt.Errorf("%w: elevation is unsupported on non-Windows v1 clients", ErrInvalidExecutionSpec) } windowsKeys := make(map[string]bool, len(spec.EnvOverrides)) for key, value := range spec.EnvOverrides { if key == "" || strings.IndexByte(key, 0) >= 0 || strings.Contains(key, "=") || strings.IndexByte(value, 0) >= 0 || !utf8.ValidString(key) || !utf8.ValidString(value) { return fmt.Errorf("%w: invalid environment override", ErrInvalidExecutionSpec) } if peerPlatform == rvboxv1.Platform_PLATFORM_WINDOWS { folded := strings.ToUpper(key) if windowsKeys[folded] { return fmt.Errorf("%w: case-colliding Windows environment keys", ErrInvalidExecutionSpec) } windowsKeys[folded] = true } } if err := validateProfiles(spec.ExecutionProfiles); err != nil { return err } switch source := spec.Source.(type) { case *rvboxv1.ExecutionSpec_CommandText: if !utf8.ValidString(source.CommandText) { return fmt.Errorf("%w: command text is not UTF-8", ErrInvalidExecutionSpec) } case *rvboxv1.ExecutionSpec_Script: if err := validateScriptDescriptor(source.Script, limits); err != nil { return err } default: return fmt.Errorf("%w: unknown source", ErrInvalidExecutionSpec) } return nil } func validateClientHello(hello *rvboxv1.ClientHello) error { if hello == nil || validateClientID(hello.ClientId) != nil || hello.SupportedProtocol == nil || hello.SentAt == nil || hello.SentAt.CheckValid() != nil || hello.ClientInstanceId == "" || len(hello.ClientInstanceId) > 128 || hello.MaxRunningCommands == 0 || hello.MaxQueuedCommands == 0 || hello.Platform == rvboxv1.Platform_PLATFORM_UNSPECIFIED || len(hello.SupportedShells) == 0 { return fmt.Errorf("%w: invalid ClientHello", ErrInvalidEnvelope) } seenShells := map[rvboxv1.ShellType]bool{} for _, shell := range hello.SupportedShells { if seenShells[shell] || !validShell(shell, hello.Platform) { return fmt.Errorf("%w: invalid ClientHello shell advertisement", ErrInvalidEnvelope) } seenShells[shell] = true } return validateProtocolRange(hello.SupportedProtocol) } func validateCommandDispatch(dispatch *rvboxv1.CommandDispatch, limits Limits, peerPlatform rvboxv1.Platform) error { if dispatch == nil || dispatch.CommandRevision == 0 || dispatch.TargetSessionGeneration == 0 || dispatch.IssueTime == nil || dispatch.IssueTime.CheckValid() != nil { return fmt.Errorf("%w: invalid command dispatch", ErrInvalidEnvelope) } if err := validateIssueUUID(dispatch.IssueUuid); err != nil { return err } if dispatch.QueueExpiryTime != nil && dispatch.QueueExpiryTime.CheckValid() != nil { return fmt.Errorf("%w: invalid queue expiry", ErrInvalidEnvelope) } return ValidateExecutionSpec(dispatch.Spec, limits, peerPlatform) } func validateCommandAccepted(accepted *rvboxv1.CommandAccepted, limits Limits) error { if accepted == nil || accepted.CommandRevision == 0 { return fmt.Errorf("%w: invalid command acceptance", ErrInvalidEnvelope) } if err := validateIssueUUID(accepted.IssueUuid); err != nil { return err } if accepted.Accepted && accepted.Rejection != nil { return fmt.Errorf("%w: accepted command carries rejection", ErrInvalidEnvelope) } if !accepted.Accepted && (accepted.Rejection == nil || accepted.Rejection.Code == rvboxv1.ControlError_CODE_UNSPECIFIED) { return fmt.Errorf("%w: rejected command lacks error", ErrInvalidEnvelope) } if accepted.Rejection != nil && (uint64(len(accepted.Rejection.Message)) > limits.MaxDetailBytes || (accepted.Rejection.IssueUuid != "" && accepted.Rejection.IssueUuid != accepted.IssueUuid)) { return fmt.Errorf("%w: invalid acceptance rejection detail", ErrInvalidEnvelope) } return nil } func validateCommandEvent(event *rvboxv1.CommandEvent, limits Limits) error { if event == nil || event.Payload == nil || event.ObservedAt == nil || event.ObservedAt.CheckValid() != nil { return fmt.Errorf("%w: invalid command event", ErrInvalidEnvelope) } if err := validateIssueAndSequence(event.IssueUuid, event.EventSeq); err != nil { return err } switch payload := event.Payload.(type) { case *rvboxv1.CommandEvent_Output: _, err := DecodeOutputChunk(payload.Output, limits.MaxRawChunkBytes) return err case *rvboxv1.CommandEvent_Lifecycle: if payload.Lifecycle == nil || payload.Lifecycle.Lifecycle == rvboxv1.CommandLifecycle_COMMAND_LIFECYCLE_UNSPECIFIED || payload.Lifecycle.CommandRevision == 0 || uint64(len(payload.Lifecycle.Detail)) > limits.MaxDetailBytes { return fmt.Errorf("%w: invalid lifecycle event", ErrInvalidEnvelope) } case *rvboxv1.CommandEvent_OutputTruncation: if payload.OutputTruncation == nil || uint64(len(payload.OutputTruncation.Reason)) > limits.MaxDetailBytes || payload.OutputTruncation.Source == rvboxv1.OutputTruncationSource_OUTPUT_TRUNCATION_SOURCE_UNSPECIFIED { return fmt.Errorf("%w: invalid output truncation", ErrInvalidEnvelope) } return domain.ValidateOutputTruncation(payload.OutputTruncation) case *rvboxv1.CommandEvent_OutputIncomplete: if payload.OutputIncomplete == nil || !validStreams(payload.OutputIncomplete.Streams) || uint64(len(payload.OutputIncomplete.Reason)) > limits.MaxDetailBytes { return fmt.Errorf("%w: invalid incomplete-output event", ErrInvalidEnvelope) } case *rvboxv1.CommandEvent_StdinAck: if payload.StdinAck == nil || payload.StdinAck.WriteSeq == 0 { return fmt.Errorf("%w: invalid stdin acknowledgement", ErrInvalidEnvelope) } case *rvboxv1.CommandEvent_SignalResult: if payload.SignalResult == nil || payload.SignalResult.CommandRevision == 0 || !validSignal(payload.SignalResult.Signal) { return fmt.Errorf("%w: invalid signal result", ErrInvalidEnvelope) } case *rvboxv1.CommandEvent_ScriptStatus: if payload.ScriptStatus == nil || payload.ScriptStatus.ReceivedBytes > limits.MaxScriptBytes { return fmt.Errorf("%w: invalid script status", ErrInvalidEnvelope) } case *rvboxv1.CommandEvent_Resource: if payload.Resource == nil || uint64(len(payload.Resource.DiagnosticDetail)) > limits.MaxDetailBytes { return fmt.Errorf("%w: missing resource snapshot", ErrInvalidEnvelope) } default: return fmt.Errorf("%w: unsupported command event payload", ErrInvalidEnvelope) } return nil } func validateScriptDescriptor(descriptor *rvboxv1.ScriptDescriptor, limits Limits) error { if descriptor == nil || descriptor.Filename == "" || len(descriptor.Filename) > 255 || descriptor.SizeBytes > limits.MaxScriptBytes || len(descriptor.Sha256) != sha256.Size || !utf8.ValidString(descriptor.Filename) { return fmt.Errorf("%w: invalid script descriptor", ErrInvalidExecutionSpec) } if filepath.Base(descriptor.Filename) != descriptor.Filename || strings.ContainsAny(descriptor.Filename, `/\`) || descriptor.Filename == "." || descriptor.Filename == ".." { return fmt.Errorf("%w: script filename must be display-only basename", ErrInvalidExecutionSpec) } if strings.TrimRight(descriptor.Filename, " .") != descriptor.Filename { return fmt.Errorf("%w: script filename has unsafe Windows suffix", ErrInvalidExecutionSpec) } for _, current := range []byte(descriptor.Filename) { if current < 0x20 || current == 0x7f { return fmt.Errorf("%w: script filename contains control byte", ErrInvalidExecutionSpec) } } base := strings.ToUpper(strings.TrimSuffix(descriptor.Filename, filepath.Ext(descriptor.Filename))) reserved := map[string]bool{"CON": true, "PRN": true, "AUX": true, "NUL": true, "CLOCK$": true} if reserved[base] || (len(base) == 4 && (strings.HasPrefix(base, "COM") || strings.HasPrefix(base, "LPT")) && base[3] >= '1' && base[3] <= '9') { return fmt.Errorf("%w: reserved script filename", ErrInvalidExecutionSpec) } return nil } func validateProfiles(profiles []rvboxv1.ExecutionProfile) error { seen := map[rvboxv1.ExecutionProfile]bool{} dimensions := map[rvboxv1.ExecutionProfile]string{ rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_MEDIUM: "cpu", rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_HEAVY: "cpu", rvboxv1.ExecutionProfile_EXECUTION_PROFILE_MEM_MEDIUM: "memory", rvboxv1.ExecutionProfile_EXECUTION_PROFILE_MEM_HEAVY: "memory", rvboxv1.ExecutionProfile_EXECUTION_PROFILE_DISK_MEDIUM: "disk", rvboxv1.ExecutionProfile_EXECUTION_PROFILE_DISK_HEAVY: "disk", } usedDimensions := map[string]bool{} for _, profile := range profiles { if profile == rvboxv1.ExecutionProfile_EXECUTION_PROFILE_UNSPECIFIED || seen[profile] { return fmt.Errorf("%w: unspecified or duplicate execution profile", ErrInvalidExecutionSpec) } seen[profile] = true if profile == rvboxv1.ExecutionProfile_EXECUTION_PROFILE_LIGHT { if len(profiles) != 1 { return fmt.Errorf("%w: LIGHT is exclusive", ErrInvalidExecutionSpec) } continue } dimension, exists := dimensions[profile] if !exists || usedDimensions[dimension] { return fmt.Errorf("%w: conflicting or unknown execution profile", ErrInvalidExecutionSpec) } usedDimensions[dimension] = true } return nil } func validateProtocolRange(value *rvboxv1.ProtocolRange) error { if value.Major == 0 || value.MinMinor > value.MaxMinor { return fmt.Errorf("%w: invalid protocol range", ErrInvalidEnvelope) } return nil } func validateProtocolVersion(value *rvboxv1.ProtocolVersion) error { if value.Major == 0 { return fmt.Errorf("%w: invalid protocol version", ErrInvalidEnvelope) } return nil } func validateIssueUUID(value string) error { if _, err := domain.ParseUUIDv7(value); err != nil { return fmt.Errorf("%w: %v", ErrInvalidEnvelope, err) } return nil } func validateIssueAndSequence(issue string, sequence uint64) error { if sequence == 0 { return fmt.Errorf("%w: sequence must be nonzero", ErrInvalidEnvelope) } return validateIssueUUID(issue) } func validateClientID(value string) error { if len(value) == 0 || len(value) > 128 { return errors.New("invalid client ID length") } for _, current := range []byte(value) { if current < 0x21 || current > 0x7e { return errors.New("client ID must be printable ASCII without spaces") } } return nil } func validShell(shell rvboxv1.ShellType, platform rvboxv1.Platform) bool { if platform == rvboxv1.Platform_PLATFORM_WINDOWS { return shell == rvboxv1.ShellType_SHELL_CMD || shell == rvboxv1.ShellType_SHELL_POWERSHELL } if platform == rvboxv1.Platform_PLATFORM_LINUX || platform == rvboxv1.Platform_PLATFORM_DARWIN || platform == rvboxv1.Platform_PLATFORM_OTHER_UNIX { return shell == rvboxv1.ShellType_SHELL_SH || shell == rvboxv1.ShellType_SHELL_BASH } return false } func validSignal(signal rvboxv1.SignalKind) bool { return signal >= rvboxv1.SignalKind_SIGNAL_HUP && signal <= rvboxv1.SignalKind_SIGNAL_USR2 } func validStreams(streams []rvboxv1.StreamKind) bool { if len(streams) == 0 { return false } seen := map[rvboxv1.StreamKind]bool{} for _, stream := range streams { if (stream != rvboxv1.StreamKind_STREAM_STDOUT && stream != rvboxv1.StreamKind_STREAM_STDERR) || seen[stream] { return false } seen[stream] = true } return true } func validateReconcileRequest(request *rvboxv1.ReconcileRequest) error { if request == nil { return fmt.Errorf("%w: missing reconcile request", ErrInvalidEnvelope) } seen := map[string]bool{} for _, target := range request.Targets { if target == nil || seen[target.IssueUuid] || target.CommandRevision == 0 || len(target.ImmutableRequestSha256) != sha256.Size { return fmt.Errorf("%w: invalid reconcile target", ErrInvalidEnvelope) } if err := validateIssueUUID(target.IssueUuid); err != nil { return err } seen[target.IssueUuid] = true } return nil } func validateReconcileSnapshot(snapshot *rvboxv1.ReconcileSnapshot) error { if snapshot == nil { return fmt.Errorf("%w: missing reconcile snapshot", ErrInvalidEnvelope) } seen := map[string]bool{} for _, command := range snapshot.RetainedCommands { if command == nil || seen[command.IssueUuid] || command.CommandRevision == 0 || command.Lifecycle == rvboxv1.CommandLifecycle_COMMAND_LIFECYCLE_UNSPECIFIED || len(command.ImmutableRequestSha256) != sha256.Size { return fmt.Errorf("%w: invalid reconcile command", ErrInvalidEnvelope) } if err := validateIssueUUID(command.IssueUuid); err != nil { return err } seen[command.IssueUuid] = true } return nil } func validateReconcileResult(result *rvboxv1.ReconcileResult) error { if result == nil { return fmt.Errorf("%w: missing reconcile result", ErrInvalidEnvelope) } seen := map[string]bool{} for _, issueUUID := range append(append([]string{}, result.TerminateLocalIssueUuids...), result.DiscardLocalTerminalIssueUuids...) { if seen[issueUUID] { return fmt.Errorf("%w: duplicate reconcile result UUID", ErrInvalidEnvelope) } if err := validateIssueUUID(issueUUID); err != nil { return err } seen[issueUUID] = true } return nil }