// Package control implements the local gRPC control plane. It is deliberately // independent of the WebSocket session implementation: all durable decisions // go through store APIs, while a later dispatcher may subscribe to queued work. package control import ( "context" "crypto/rand" "crypto/sha256" "encoding/binary" "errors" "fmt" "math" "time" rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/agentproto" "github.com/rvbox/rvbox/internal/domain" "github.com/rvbox/rvbox/internal/server/store" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" "google.golang.org/protobuf/proto" "google.golang.org/protobuf/types/known/durationpb" "google.golang.org/protobuf/types/known/timestamppb" ) const ( defaultPageSize = 100 maxPageSize = 1000 defaultQueueTTL = 15 * time.Minute clientsCursorVersion = "clients-v1" ) // Options controls the non-transport dependencies of Service. CursorKey is // process-local by default, which invalidates old cursors after a restart // rather than accepting a token whose query snapshot is no longer meaningful. type Options struct { Store *store.Store DefaultQueueTTL time.Duration // DefaultQueueTTLSet distinguishes an explicitly configured zero (the // documented indefinite-queue setting) from an omitted option in tests or // embedders that want the compiled 15-minute default. DefaultQueueTTLSet bool Limits agentproto.Limits Now func() time.Time CursorKey []byte WakeClient func(string) bool } type Service struct { rvboxv1.UnimplementedControlServer store *store.Store defaultQueueTTL time.Duration limits agentproto.Limits now func() time.Time cursors *domain.CursorCodec wakeClient func(string) bool } func NewService(options Options) (*Service, error) { if options.Store == nil { return nil, errors.New("control service requires a store") } if options.DefaultQueueTTL < 0 { return nil, errors.New("default queue TTL must be non-negative") } if !options.DefaultQueueTTLSet && options.DefaultQueueTTL == 0 { options.DefaultQueueTTL = defaultQueueTTL } if options.Limits.MaxEnvelopeBytes == 0 { options.Limits = agentproto.DefaultLimits() } if options.Now == nil { options.Now = time.Now } key := append([]byte(nil), options.CursorKey...) if len(key) == 0 { key = make([]byte, sha256.Size) if _, err := rand.Read(key); err != nil { return nil, fmt.Errorf("generate cursor key: %w", err) } } codec, err := domain.NewCursorCodec(key) if err != nil { return nil, err } return &Service{store: options.Store, defaultQueueTTL: options.DefaultQueueTTL, limits: options.Limits, now: options.Now, cursors: codec, wakeClient: options.WakeClient}, nil } func (service *Service) ListClients(ctx context.Context, request *rvboxv1.ListClientsRequest) (*rvboxv1.ListClientsResponse, error) { if request == nil { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "request is required") } pageSize, err := pageSize(request.GetPageSize()) if err != nil { return nil, err } filter := domain.HashCursorFilters([]byte(clientsCursorVersion)) after := "" if request.GetPageToken() != "" { cursor, decodeErr := service.cursors.Decode(request.GetPageToken(), domain.CursorKindClients, filter) if decodeErr != nil || string(cursor.SnapshotBoundary) != clientsCursorVersion { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "invalid client page token") } after = string(cursor.Position) } clients, hasNext, err := service.store.ListClientViews(ctx, after, pageSize) if err != nil { return nil, mapStoreError(err) } response := &rvboxv1.ListClientsResponse{Clients: make([]*rvboxv1.ClientSummary, 0, len(clients))} for _, client := range clients { converted, convertErr := clientSummary(client) if convertErr != nil { return nil, mapStoreError(convertErr) } response.Clients = append(response.Clients, converted) } if hasNext && len(clients) > 0 { response.NextPageToken, err = service.cursors.Encode(domain.Cursor{Kind: domain.CursorKindClients, FilterHash: filter, Position: []byte(clients[len(clients)-1].ClientID), SnapshotBoundary: []byte(clientsCursorVersion)}) if err != nil { return nil, mapStoreError(err) } } return response, nil } func (service *Service) GetClient(ctx context.Context, request *rvboxv1.GetClientRequest) (*rvboxv1.GetClientResponse, error) { if request == nil || request.GetClientId() == "" { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id is required") } client, err := service.store.GetClientView(ctx, request.GetClientId()) if err != nil { return nil, mapStoreError(err) } summary, err := clientSummary(client) if err != nil { return nil, mapStoreError(err) } return &rvboxv1.GetClientResponse{Client: summary}, nil } func (service *Service) ListCommands(ctx context.Context, request *rvboxv1.ListCommandsRequest) (*rvboxv1.ListCommandsResponse, error) { if request == nil || request.GetClientId() == "" { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id is required") } if _, err := service.store.GetClientView(ctx, request.GetClientId()); err != nil { return nil, mapStoreError(err) } limit, err := pageSize(request.GetPageSize()) if err != nil { return nil, err } filterBytes := []byte(fmt.Sprintf("commands\x00%s\x00%t", request.GetClientId(), request.GetIncludeTerminal())) filter := domain.HashCursorFilters(filterBytes) page := store.CommandPage{ClientID: request.GetClientId(), IncludeTerminal: request.GetIncludeTerminal(), Limit: limit, SnapshotBoundary: service.now().UnixNano()} if request.GetPageToken() != "" { cursor, decodeErr := service.cursors.Decode(request.GetPageToken(), domain.CursorKindCommands, filter) if decodeErr != nil || len(cursor.Position) != 24 || len(cursor.SnapshotBoundary) != 8 { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "invalid command page token") } page.SnapshotBoundary = int64(binary.BigEndian.Uint64(cursor.SnapshotBoundary)) page.AfterTime = int64(binary.BigEndian.Uint64(cursor.Position[:8])) copy(page.AfterUUID[:], cursor.Position[8:]) page.HasAfter = true } commands, hasNext, err := service.store.ListCommandViews(ctx, page) if err != nil { return nil, mapStoreError(err) } response := &rvboxv1.ListCommandsResponse{Commands: make([]*rvboxv1.CommandRecord, 0, len(commands))} for _, command := range commands { converted, convertErr := commandRecord(command) if convertErr != nil { return nil, mapStoreError(convertErr) } response.Commands = append(response.Commands, converted) } if hasNext && len(commands) > 0 { position := make([]byte, 24) binary.BigEndian.PutUint64(position[:8], uint64(commands[len(commands)-1].IssueTime.UnixNano())) copy(position[8:], commands[len(commands)-1].IssueUUID[:]) boundary := make([]byte, 8) binary.BigEndian.PutUint64(boundary, uint64(page.SnapshotBoundary)) response.NextPageToken, err = service.cursors.Encode(domain.Cursor{Kind: domain.CursorKindCommands, FilterHash: filter, Position: position, SnapshotBoundary: boundary}) if err != nil { return nil, mapStoreError(err) } } return response, nil } func (service *Service) GetCommand(ctx context.Context, request *rvboxv1.GetCommandRequest) (*rvboxv1.GetCommandResponse, error) { if request == nil || request.GetIssueUuid() == "" { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "issue_uuid is required") } issue, err := domain.ParseUUIDv7(request.GetIssueUuid()) if err != nil { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "issue_uuid must be canonical UUIDv7") } command, err := service.store.GetCommandView(ctx, request.GetClientId(), issue) if err != nil { return nil, mapStoreError(err) } record, err := commandRecord(command) if err != nil { return nil, mapStoreError(err) } return &rvboxv1.GetCommandResponse{Command: record}, nil } // RunCommand durably admits command-text work. Script payload persistence and // chunk dispatch are intentionally kept behind the same immutable boundary and // are rejected until their command_payloads path is wired to the dispatcher. func (service *Service) RunCommand(ctx context.Context, request *rvboxv1.RunCommandRequest) (*rvboxv1.RunCommandResponse, error) { if request == nil || request.GetTargetClientId() == "" || request.GetSpec() == nil { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "target_client_id and spec are required") } client, err := service.store.GetClientView(ctx, request.GetTargetClientId()) if err != nil { return nil, mapStoreError(err) } platform := rvboxv1.Platform(client.Platform) spec := proto.Clone(request.GetSpec()).(*rvboxv1.ExecutionSpec) if spec.GetShellType() == rvboxv1.ShellType_SHELL_TYPE_UNSPECIFIED { switch platform { case rvboxv1.Platform_PLATFORM_WINDOWS: spec.ShellType = rvboxv1.ShellType_SHELL_POWERSHELL case rvboxv1.Platform_PLATFORM_LINUX, rvboxv1.Platform_PLATFORM_DARWIN, rvboxv1.Platform_PLATFORM_OTHER_UNIX: spec.ShellType = rvboxv1.ShellType_SHELL_SH default: return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "target client platform is unspecified") } } if _, script := spec.Source.(*rvboxv1.ExecutionSpec_Script); script { return nil, status.Error(codes.Unimplemented, "script command admission is not enabled until payload dispatch is implemented") } if len(request.GetScriptContent()) != 0 { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "script_content requires a script source") } if err := agentproto.ValidateExecutionSpec(spec, service.limits, platform); err != nil { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, err.Error()) } if !advertisedShell(client.SupportedShells, spec.GetShellType()) { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "target client does not advertise the requested shell") } issue, err := service.requestID(request.GetRequestId()) if err != nil { return nil, err } now := service.now().UTC() expiry, err := service.queueExpiry(request.GetQueueTtl(), now) if err != nil { return nil, err } canonical := &rvboxv1.RunCommandRequest{TargetClientId: request.GetTargetClientId(), Spec: spec, QueueTtl: request.GetQueueTtl()} encoded, err := proto.MarshalOptions{Deterministic: true}.Marshal(canonical) if err != nil { return nil, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "canonicalize command request") } hash := sha256.Sum256(encoded) queued, err := service.store.QueueCommand(ctx, store.QueueCommandInput{ IssueUUID: issue, ClientID: request.GetTargetClientId(), IssueTime: now, ReceiptTime: now, QueueExpiryTime: expiry, ImmutableSHA256: hash, ExecutionSpec: mustMarshal(spec), }) if err != nil { return nil, mapStoreError(err) } if service.wakeClient != nil { _ = service.wakeClient(request.GetTargetClientId()) } return &rvboxv1.RunCommandResponse{IssueUuid: issue.String(), Lifecycle: rvboxv1.CommandLifecycle(queued.Lifecycle)}, nil } func (service *Service) requestID(value string) (domain.UUID, error) { if value == "" { issue, err := domain.NewUUIDv7() if err != nil { return domain.UUID{}, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "generate request ID") } return issue, nil } issue, err := domain.ParseUUIDv7(value) if err != nil { return domain.UUID{}, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "request_id must be canonical UUIDv7") } return issue, nil } func (service *Service) queueExpiry(value *durationpb.Duration, now time.Time) (*time.Time, error) { if value == nil { expiry := now.Add(service.defaultQueueTTL) return &expiry, nil } if err := value.CheckValid(); err != nil || value.AsDuration() < 0 { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "queue_ttl must be non-negative") } if value.AsDuration() == 0 { return nil, nil } duration := value.AsDuration() if now.UnixNano() > math.MaxInt64-int64(duration) { return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "queue_ttl overflows server clock") } expiry := now.Add(duration) return &expiry, nil } func mustMarshal(message proto.Message) []byte { encoded, _ := (proto.MarshalOptions{Deterministic: true}).Marshal(message) return encoded } func advertisedShell(encoded []byte, shell rvboxv1.ShellType) bool { var hello rvboxv1.ClientHello if err := proto.Unmarshal(encoded, &hello); err == nil { for _, advertised := range hello.GetSupportedShells() { if advertised == shell { return true } } } return false } func clientSummary(view store.ClientView) (*rvboxv1.ClientSummary, error) { result := &rvboxv1.ClientSummary{ClientId: view.ClientID, Connected: view.Connected, Platform: rvboxv1.Platform(view.Platform), Architecture: view.Architecture, DaemonVersion: view.DaemonVersion, DaemonCwd: view.DaemonCWD, RunningCommands: view.RunningCommands, QueuedCommands: view.QueuedCommands, ClientInstanceId: uuidString(view.ClientInstanceID)} var hello rvboxv1.ClientHello if err := proto.Unmarshal(view.SupportedShells, &hello); err != nil { return nil, fmt.Errorf("decode stored shell advertisement: %w", err) } result.SupportedShells = append([]rvboxv1.ShellType(nil), hello.GetSupportedShells()...) if view.ConnectedAt != nil { result.ConnectedAt = timestamppb.New(*view.ConnectedAt) } if view.LastSeenAt != nil { result.LastSeenAt = timestamppb.New(*view.LastSeenAt) } if view.PendingInstanceID != nil { result.PendingInstanceId = uuidString(*view.PendingInstanceID) } if view.PendingInstanceSeenAt != nil { result.PendingInstanceSeenAt = timestamppb.New(*view.PendingInstanceSeenAt) } return result, nil } func commandRecord(view store.CommandView) (*rvboxv1.CommandRecord, error) { spec := &rvboxv1.ExecutionSpec{} if err := proto.Unmarshal(view.ExecutionSpec, spec); err != nil { return nil, fmt.Errorf("decode stored execution spec: %w", err) } result := &rvboxv1.CommandRecord{IssueUuid: view.IssueUUID.String(), TargetClientId: view.ClientID, IssueTime: timestamppb.New(view.IssueTime), ServerReceiptTime: timestamppb.New(view.ServerReceiptTime), Spec: spec, Lifecycle: rvboxv1.CommandLifecycle(view.Lifecycle), LastEventSeq: view.LastEventSeq, OutputTruncated: view.OutputTruncated, OutputIncomplete: view.OutputIncomplete, RetainedCompressedBytes: view.RetainedCompressedBytes, CommandRevision: view.Revision} if view.QueueExpiryTime != nil { result.QueueExpiryTime = timestamppb.New(*view.QueueExpiryTime) } if view.TerminalTime != nil { result.TerminalTime = timestamppb.New(*view.TerminalTime) } if view.ExitCode != nil { result.ExitCode = view.ExitCode } if len(view.WindowsIdentity) > 0 { identity := &rvboxv1.WindowsExecutionIdentity{} if err := proto.Unmarshal(view.WindowsIdentity, identity); err != nil { return nil, fmt.Errorf("decode stored Windows identity: %w", err) } result.WindowsExecutionIdentity = identity } return result, nil } func uuidString(value [16]byte) string { return domain.UUID(value).String() } func pageSize(value uint32) (uint32, error) { if value == 0 { return defaultPageSize, nil } if value > maxPageSize { return 0, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "page_size exceeds the maximum") } return value, nil } func controlError(code codes.Code, rvCode rvboxv1.ControlError_Code, message string) error { detail := &rvboxv1.ControlError{Code: rvCode, Message: message, Retryable: rvCode == rvboxv1.ControlError_OFFLINE || rvCode == rvboxv1.ControlError_CAPACITY_EXHAUSTED || rvCode == rvboxv1.ControlError_TRANSIENT || rvCode == rvboxv1.ControlError_INTERNAL} result := status.New(code, message) withDetails, err := result.WithDetails(detail) if err == nil { return withDetails.Err() } return result.Err() } func mapStoreError(err error) error { if err == nil { return nil } switch { case errors.Is(err, store.ErrClientNotFound), errors.Is(err, store.ErrCommandNotFound), errors.Is(err, store.ErrIncidentNotFound): return controlError(codes.NotFound, rvboxv1.ControlError_NOT_FOUND, err.Error()) case errors.Is(err, store.ErrCommandConflict), errors.Is(err, store.ErrMutationConflict): return controlError(codes.AlreadyExists, rvboxv1.ControlError_CONFLICT, err.Error()) case errors.Is(err, store.ErrCapacityExhausted): return controlError(codes.ResourceExhausted, rvboxv1.ControlError_CAPACITY_EXHAUSTED, err.Error()) case errors.Is(err, store.ErrStoreClosed): return controlError(codes.Unavailable, rvboxv1.ControlError_OFFLINE, err.Error()) default: return controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, err.Error()) } }