847 lines
37 KiB
Go
847 lines
37 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"
|
|
"sort"
|
|
"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
|
|
TakeoverTTL time.Duration
|
|
Limits agentproto.Limits
|
|
Now func() time.Time
|
|
CursorKey []byte
|
|
WakeClient func(string) bool
|
|
}
|
|
|
|
type Service struct {
|
|
rvboxv1.UnimplementedControlServer
|
|
store *store.Store
|
|
defaultQueueTTL time.Duration
|
|
takeoverTTL 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.TakeoverTTL < 0 {
|
|
return nil, errors.New("takeover TTL must be non-negative")
|
|
}
|
|
if options.TakeoverTTL == 0 {
|
|
options.TakeoverTTL = 5 * time.Minute
|
|
}
|
|
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, takeoverTTL: options.TakeoverTTL, 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
|
|
}
|
|
|
|
// FollowCommand streams the durable event history and then waits for newly
|
|
// committed events. It polls the store rather than maintaining a second event
|
|
// bus, so a reconnect can resume from the last event sequence without losing a
|
|
// commit that raced the stream cancellation.
|
|
func (service *Service) FollowCommand(request *rvboxv1.FollowCommandRequest, stream rvboxv1.Control_FollowCommandServer) error {
|
|
if request == nil || request.GetIssueUuid() == "" || request.GetClientId() == "" {
|
|
return controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id and issue_uuid are required")
|
|
}
|
|
issue, err := domain.ParseUUIDv7(request.GetIssueUuid())
|
|
if err != nil {
|
|
return controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "issue_uuid must be canonical UUIDv7")
|
|
}
|
|
view, err := service.store.GetCommandView(stream.Context(), request.GetClientId(), issue)
|
|
if err != nil {
|
|
return mapStoreError(err)
|
|
}
|
|
after := request.GetAfterEventSeq()
|
|
if !request.GetIncludeExisting() && after == 0 {
|
|
after = view.LastEventSeq
|
|
}
|
|
for {
|
|
events, readErr := service.store.ReadCommandEvents(stream.Context(), issue, after, 1000)
|
|
if readErr != nil {
|
|
return mapStoreError(readErr)
|
|
}
|
|
for _, stored := range events {
|
|
payload, decodeErr := store.DecodeEventPayload(stored)
|
|
if decodeErr != nil {
|
|
return mapStoreError(decodeErr)
|
|
}
|
|
event := &rvboxv1.CommandEvent{}
|
|
if unmarshalErr := proto.Unmarshal(payload, event); unmarshalErr != nil || event.GetIssueUuid() != issue.String() || event.GetEventSeq() != stored.EventSeq || len(event.GetImmutableEventSha256()) != sha256.Size {
|
|
return controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "stored command event is invalid")
|
|
}
|
|
digest, digestErr := agentproto.CommandEventDigest(event)
|
|
if digestErr != nil || string(digest[:]) != string(event.GetImmutableEventSha256()) {
|
|
return controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "stored command event digest is invalid")
|
|
}
|
|
if err := stream.Send(&rvboxv1.FollowCommandResponse{Item: &rvboxv1.FollowCommandResponse_Event{Event: event}, ServerReceiptTime: timestamppb.New(time.Unix(0, stored.ReceiptUnixNano).UTC())}); err != nil {
|
|
return err
|
|
}
|
|
after = stored.EventSeq
|
|
}
|
|
view, err = service.store.GetCommandView(stream.Context(), request.GetClientId(), issue)
|
|
if err != nil {
|
|
return mapStoreError(err)
|
|
}
|
|
if domain.IsTerminal(rvboxv1.CommandLifecycle(view.Lifecycle)) && after >= view.LastEventSeq {
|
|
return nil
|
|
}
|
|
timer := time.NewTimer(100 * time.Millisecond)
|
|
select {
|
|
case <-stream.Context().Done():
|
|
timer.Stop()
|
|
return stream.Context().Err()
|
|
case <-timer.C:
|
|
}
|
|
}
|
|
}
|
|
|
|
// GetOutput returns byte-bounded, cursor-resumable slices from output events.
|
|
// The cursor binds the selected streams and a command-event snapshot boundary;
|
|
// concurrent appends therefore cannot reorder or duplicate earlier bytes.
|
|
func (service *Service) GetOutput(ctx context.Context, request *rvboxv1.GetOutputRequest) (*rvboxv1.GetOutputResponse, error) {
|
|
if request == nil || request.GetClientId() == "" || request.GetIssueUuid() == "" {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id and issue_uuid are 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")
|
|
}
|
|
if _, err := service.store.GetClientView(ctx, request.GetClientId()); err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
streams, err := normalizeStreams(request.GetStreams())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
maxBytes := request.GetMaxBytes()
|
|
if maxBytes == 0 {
|
|
maxBytes = 1 << 20
|
|
}
|
|
if maxBytes > 16<<20 {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "max_bytes exceeds the control limit")
|
|
}
|
|
filterBytes := []byte(fmt.Sprintf("output\x00%s\x00%v", request.GetClientId(), streams))
|
|
filter := domain.HashCursorFilters(filterBytes)
|
|
var eventSeq, offset, boundary uint64
|
|
if request.GetPageToken() != "" {
|
|
cursor, decodeErr := service.cursors.Decode(request.GetPageToken(), domain.CursorKindOutput, filter)
|
|
if decodeErr != nil || len(cursor.Position) != 16 || len(cursor.SnapshotBoundary) != 8 {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "invalid output page token")
|
|
}
|
|
eventSeq = binary.BigEndian.Uint64(cursor.Position[:8])
|
|
offset = binary.BigEndian.Uint64(cursor.Position[8:])
|
|
boundary = binary.BigEndian.Uint64(cursor.SnapshotBoundary)
|
|
} else {
|
|
view, viewErr := service.store.GetCommandView(ctx, request.GetClientId(), issue)
|
|
if viewErr != nil {
|
|
return nil, mapStoreError(viewErr)
|
|
}
|
|
boundary = view.LastEventSeq
|
|
eventSeq = request.GetAfterEventSeq()
|
|
}
|
|
queryAfter := eventSeq
|
|
if offset > 0 && eventSeq > 0 {
|
|
queryAfter = eventSeq - 1
|
|
}
|
|
storedEvents, err := service.store.ReadCommandEvents(ctx, issue, queryAfter, 1000)
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
response := &rvboxv1.GetOutputResponse{}
|
|
remaining := maxBytes
|
|
hasMore := false
|
|
for _, stored := range storedEvents {
|
|
if stored.EventSeq > boundary {
|
|
break
|
|
}
|
|
payload, decodeErr := store.DecodeEventPayload(stored)
|
|
if decodeErr != nil {
|
|
return nil, mapStoreError(decodeErr)
|
|
}
|
|
event := &rvboxv1.CommandEvent{}
|
|
if unmarshalErr := proto.Unmarshal(payload, event); unmarshalErr != nil || event.GetIssueUuid() != issue.String() || event.GetEventSeq() != stored.EventSeq {
|
|
return nil, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "stored command event is invalid")
|
|
}
|
|
output := event.GetOutput()
|
|
if output == nil || !containsStream(streams, output.GetStream()) {
|
|
if stored.EventSeq == eventSeq {
|
|
offset = 0
|
|
}
|
|
continue
|
|
}
|
|
data, decodeErr := agentproto.DecodeOutputChunk(output, service.limits.MaxRawChunkBytes)
|
|
if decodeErr != nil {
|
|
return nil, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "stored output chunk is invalid")
|
|
}
|
|
start := uint64(0)
|
|
if stored.EventSeq == eventSeq {
|
|
start = offset
|
|
}
|
|
if start > uint64(len(data)) {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "output cursor is outside its event")
|
|
}
|
|
if start == uint64(len(data)) {
|
|
eventSeq, offset = stored.EventSeq, 0
|
|
continue
|
|
}
|
|
if remaining == 0 {
|
|
hasMore = true
|
|
break
|
|
}
|
|
end := uint64(len(data))
|
|
if end-start > remaining {
|
|
end = start + remaining
|
|
hasMore = true
|
|
}
|
|
response.Output = append(response.Output, &rvboxv1.OutputSlice{EventSeq: stored.EventSeq, ObservedAt: timestamppb.New(time.Unix(0, stored.ObservedUnixNano).UTC()), Stream: output.GetStream(), Data: append([]byte(nil), data[start:end]...), EventByteOffset: start, EndOfEvent: end == uint64(len(data)), ServerReceiptTime: timestamppb.New(time.Unix(0, stored.ReceiptUnixNano).UTC())})
|
|
remaining -= end - start
|
|
eventSeq, offset = stored.EventSeq, end
|
|
if end < uint64(len(data)) {
|
|
break
|
|
}
|
|
if remaining == 0 {
|
|
hasMore = true
|
|
break
|
|
}
|
|
}
|
|
if len(storedEvents) == 1000 {
|
|
hasMore = true
|
|
}
|
|
view, viewErr := service.store.GetCommandView(ctx, request.GetClientId(), issue)
|
|
if viewErr == nil {
|
|
response.OutputTruncated = view.OutputTruncated
|
|
if view.OutputIncomplete {
|
|
response.Incomplete = &rvboxv1.OutputIncomplete{Reason: "persisted output capture is incomplete"}
|
|
}
|
|
}
|
|
if hasMore {
|
|
position := make([]byte, 16)
|
|
binary.BigEndian.PutUint64(position[:8], eventSeq)
|
|
binary.BigEndian.PutUint64(position[8:], offset)
|
|
snapshot := make([]byte, 8)
|
|
binary.BigEndian.PutUint64(snapshot, boundary)
|
|
response.NextPageToken, err = service.cursors.Encode(domain.Cursor{Kind: domain.CursorKindOutput, FilterHash: filter, Position: position, SnapshotBoundary: snapshot})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
}
|
|
return response, nil
|
|
}
|
|
|
|
func (service *Service) AppendStdin(ctx context.Context, request *rvboxv1.AppendStdinRequest) (*rvboxv1.AppendStdinResponse, error) {
|
|
if request == nil || request.GetClientId() == "" || request.GetIssueUuid() == "" || len(request.GetData()) == 0 {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id, issue_uuid, and non-empty data are required")
|
|
}
|
|
if uint64(len(request.GetData())) > service.limits.MaxRawChunkBytes {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "stdin data exceeds the raw chunk limit")
|
|
}
|
|
issue, err := domain.ParseUUIDv7(request.GetIssueUuid())
|
|
if err != nil {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "issue_uuid must be canonical UUIDv7")
|
|
}
|
|
requestUUID, err := service.requestID(request.GetRequestId())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
canonical := proto.Clone(request).(*rvboxv1.AppendStdinRequest)
|
|
canonical.RequestId = ""
|
|
encoded, err := proto.MarshalOptions{Deterministic: true}.Marshal(canonical)
|
|
if err != nil {
|
|
return nil, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "canonicalize stdin request")
|
|
}
|
|
result, err := service.store.AppendStdin(ctx, store.StdinWriteInput{IssueUUID: issue, ClientID: request.GetClientId(), RequestUUID: requestUUID, Data: request.GetData(), AppendNewline: request.GetAppendNewline(), ImmutableHash: sha256.Sum256(encoded), OccurredAt: service.now()})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
if service.wakeClient != nil {
|
|
_ = service.wakeClient(request.GetClientId())
|
|
}
|
|
return &rvboxv1.AppendStdinResponse{WriteSeq: result.WriteSeq}, nil
|
|
}
|
|
|
|
func (service *Service) CloseStdin(ctx context.Context, request *rvboxv1.CloseStdinRequest) (*rvboxv1.CloseStdinResponse, error) {
|
|
if request == nil || request.GetClientId() == "" || request.GetIssueUuid() == "" {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id and issue_uuid are 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")
|
|
}
|
|
requestUUID, err := service.requestID(request.GetRequestId())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
canonical := proto.Clone(request).(*rvboxv1.CloseStdinRequest)
|
|
canonical.RequestId = ""
|
|
encoded, err := proto.MarshalOptions{Deterministic: true}.Marshal(canonical)
|
|
if err != nil {
|
|
return nil, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "canonicalize close-stdin request")
|
|
}
|
|
result, err := service.store.CloseStdin(ctx, store.StdinWriteInput{IssueUUID: issue, ClientID: request.GetClientId(), RequestUUID: requestUUID, ImmutableHash: sha256.Sum256(encoded), OccurredAt: service.now()})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
if service.wakeClient != nil {
|
|
_ = service.wakeClient(request.GetClientId())
|
|
}
|
|
return &rvboxv1.CloseStdinResponse{WriteSeq: result.WriteSeq}, nil
|
|
}
|
|
|
|
func (service *Service) SignalCommand(ctx context.Context, request *rvboxv1.ControlSignalCommandRequest) (*rvboxv1.ControlSignalCommandResponse, error) {
|
|
if request == nil || request.GetClientId() == "" || request.GetIssueUuid() == "" || request.GetSignal() < rvboxv1.SignalKind_SIGNAL_HUP || request.GetSignal() > rvboxv1.SignalKind_SIGNAL_USR2 {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id, issue_uuid, and a portable signal are 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")
|
|
}
|
|
requestUUID, err := service.requestID(request.GetRequestId())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
canonical := proto.Clone(request).(*rvboxv1.ControlSignalCommandRequest)
|
|
canonical.RequestId = ""
|
|
encoded, err := proto.MarshalOptions{Deterministic: true}.Marshal(canonical)
|
|
if err != nil {
|
|
return nil, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "canonicalize signal request")
|
|
}
|
|
result, err := service.store.SignalCommand(ctx, store.SignalInput{IssueUUID: issue, ClientID: request.GetClientId(), RequestUUID: requestUUID, Signal: request.GetSignal(), Hash: sha256.Sum256(encoded), OccurredAt: service.now()})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
return &rvboxv1.ControlSignalCommandResponse{CommandRevision: result.CommandRevision}, nil
|
|
}
|
|
|
|
func (service *Service) ListStorageIncidents(ctx context.Context, request *rvboxv1.ListStorageIncidentsRequest) (*rvboxv1.ListStorageIncidentsResponse, error) {
|
|
if request == nil {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "request is required")
|
|
}
|
|
limit, err := pageSize(request.GetPageSize())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
filterBytes := []byte(fmt.Sprintf("incidents\x00%t", request.GetIncludeResolved()))
|
|
filter := domain.HashCursorFilters(filterBytes)
|
|
page := store.IncidentPage{IncludeResolved: request.GetIncludeResolved(), Limit: limit, SnapshotBoundary: service.now().UnixNano()}
|
|
if request.GetPageToken() != "" {
|
|
cursor, decodeErr := service.cursors.Decode(request.GetPageToken(), domain.CursorKindIncidents, filter)
|
|
if decodeErr != nil || len(cursor.Position) != 24 || len(cursor.SnapshotBoundary) != 8 {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "invalid incident page token")
|
|
}
|
|
page.AfterDetectedAt = int64(binary.BigEndian.Uint64(cursor.Position[:8]))
|
|
copy(page.AfterUUID[:], cursor.Position[8:])
|
|
page.SnapshotBoundary = int64(binary.BigEndian.Uint64(cursor.SnapshotBoundary))
|
|
page.HasAfter = true
|
|
}
|
|
incidents, hasNext, err := service.store.ListIncidentViews(ctx, page)
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
response := &rvboxv1.ListStorageIncidentsResponse{Incidents: make([]*rvboxv1.StorageIncident, 0, len(incidents))}
|
|
for _, incident := range incidents {
|
|
response.Incidents = append(response.Incidents, incidentRecord(incident))
|
|
}
|
|
if hasNext && len(incidents) > 0 {
|
|
position := make([]byte, 24)
|
|
binary.BigEndian.PutUint64(position[:8], uint64(incidents[len(incidents)-1].DetectedAt.UnixNano()))
|
|
copy(position[8:], incidents[len(incidents)-1].IncidentUUID[:])
|
|
snapshot := make([]byte, 8)
|
|
binary.BigEndian.PutUint64(snapshot, uint64(page.SnapshotBoundary))
|
|
response.NextPageToken, err = service.cursors.Encode(domain.Cursor{Kind: domain.CursorKindIncidents, FilterHash: filter, Position: position, SnapshotBoundary: snapshot})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
}
|
|
return response, nil
|
|
}
|
|
|
|
func (service *Service) RepairStorageIncident(ctx context.Context, request *rvboxv1.RepairStorageIncidentRequest) (*rvboxv1.RepairStorageIncidentResponse, error) {
|
|
if request == nil || request.GetIncidentId() == "" {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "incident_id is required")
|
|
}
|
|
incidentID, err := domain.ParseUUIDv7(request.GetIncidentId())
|
|
if err != nil {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "incident_id must be canonical UUIDv7")
|
|
}
|
|
requestID, err := service.requestID(request.GetRequestId())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
resolved, err := service.store.ResolveIncident(ctx, store.IncidentResolution{RequestUUID: [16]byte(requestID), IncidentUUID: [16]byte(incidentID), State: store.IncidentRepaired, ResolvedAt: service.now()})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
view, err := service.store.GetIncidentView(ctx, resolved.IncidentUUID)
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
return &rvboxv1.RepairStorageIncidentResponse{Incident: incidentRecord(view)}, nil
|
|
}
|
|
|
|
func (service *Service) AcknowledgeStorageIncident(ctx context.Context, request *rvboxv1.AcknowledgeStorageIncidentRequest) (*rvboxv1.AcknowledgeStorageIncidentResponse, error) {
|
|
if request == nil || request.GetIncidentId() == "" || request.GetNote() == "" {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "incident_id and note are required")
|
|
}
|
|
incidentID, err := domain.ParseUUIDv7(request.GetIncidentId())
|
|
if err != nil {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "incident_id must be canonical UUIDv7")
|
|
}
|
|
requestID, err := service.requestID(request.GetRequestId())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
resolved, err := service.store.ResolveIncident(ctx, store.IncidentResolution{RequestUUID: [16]byte(requestID), IncidentUUID: [16]byte(incidentID), State: store.IncidentAcknowledged, ResolvedAt: service.now(), Note: request.GetNote()})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
view, err := service.store.GetIncidentView(ctx, resolved.IncidentUUID)
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
return &rvboxv1.AcknowledgeStorageIncidentResponse{Incident: incidentRecord(view)}, nil
|
|
}
|
|
|
|
func (service *Service) AuthorizeClientTakeover(ctx context.Context, request *rvboxv1.AuthorizeClientTakeoverRequest) (*rvboxv1.AuthorizeClientTakeoverResponse, error) {
|
|
if request == nil || request.GetClientId() == "" || request.GetClientInstanceId() == "" {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id and client_instance_id are required")
|
|
}
|
|
instance, err := domain.ParseUUIDv7(request.GetClientInstanceId())
|
|
if err != nil {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_instance_id must be canonical UUIDv7")
|
|
}
|
|
requestID, err := service.requestID(request.GetRequestId())
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
expires := service.now().UTC().Add(service.takeoverTTL)
|
|
value, _, err := service.store.AuthorizeClientTakeover(ctx, store.TakeoverAuthorization{ClientID: request.GetClientId(), ClientInstanceID: [16]byte(instance), RequestID: [16]byte(requestID), ExpiresAt: expires})
|
|
if err != nil {
|
|
return nil, mapStoreError(err)
|
|
}
|
|
return &rvboxv1.AuthorizeClientTakeoverResponse{ExpiresAt: timestamppb.New(value)}, nil
|
|
}
|
|
|
|
func incidentRecord(view store.IncidentView) *rvboxv1.StorageIncident {
|
|
result := &rvboxv1.StorageIncident{IncidentId: domain.UUID(view.IncidentUUID).String(), DetectedAt: timestamppb.New(view.DetectedAt), State: rvboxv1.StorageIncidentState(view.State), Scope: string(view.Scope), ClientId: view.ClientID, Summary: view.Summary, DataLoss: view.DataLoss, AutomaticallyRepairable: view.AutomaticallyRepairable}
|
|
if view.ResolvedAt != nil {
|
|
result.ResolvedAt = timestamppb.New(*view.ResolvedAt)
|
|
}
|
|
if view.IssueUUID != nil {
|
|
result.IssueUuid = domain.UUID(*view.IssueUUID).String()
|
|
}
|
|
return result
|
|
}
|
|
|
|
func normalizeStreams(streams []rvboxv1.StreamKind) ([]rvboxv1.StreamKind, error) {
|
|
if len(streams) == 0 {
|
|
return []rvboxv1.StreamKind{rvboxv1.StreamKind_STREAM_STDOUT, rvboxv1.StreamKind_STREAM_STDERR}, nil
|
|
}
|
|
seen := make(map[rvboxv1.StreamKind]bool, len(streams))
|
|
result := append([]rvboxv1.StreamKind(nil), streams...)
|
|
for _, stream := range result {
|
|
if stream != rvboxv1.StreamKind_STREAM_STDOUT && stream != rvboxv1.StreamKind_STREAM_STDERR || seen[stream] {
|
|
return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "streams must contain unique stdout/stderr values")
|
|
}
|
|
seen[stream] = true
|
|
}
|
|
sort.Slice(result, func(left, right int) bool { return result[left] < result[right] })
|
|
return result, nil
|
|
}
|
|
|
|
func containsStream(streams []rvboxv1.StreamKind, value rvboxv1.StreamKind) bool {
|
|
for _, stream := range streams {
|
|
if stream == value {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
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.ErrCommandTerminal), errors.Is(err, store.ErrSignalDeliveryUnavailable):
|
|
return controlError(codes.FailedPrecondition, rvboxv1.ControlError_UNSUPPORTED, err.Error())
|
|
case errors.Is(err, store.ErrTakeoverRequired), errors.Is(err, store.ErrTakeoverMismatch), errors.Is(err, store.ErrTakeoverAlreadyGranted):
|
|
return controlError(codes.FailedPrecondition, rvboxv1.ControlError_CONFLICT, 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())
|
|
}
|
|
}
|