Files
rvbox/internal/server/control/service.go
T

407 lines
17 KiB
Go

// 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
}
type Service struct {
rvboxv1.UnimplementedControlServer
store *store.Store
defaultQueueTTL time.Duration
limits agentproto.Limits
now func() time.Time
cursors *domain.CursorCodec
}
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}, 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)
}
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())
}
}