feat: implement control history and mutation APIs
This commit is contained in:
@@ -71,7 +71,7 @@ func run(configPath string) error {
|
||||
defer func() { _ = cleanupControl() }()
|
||||
registry := session.NewRegistry()
|
||||
controlService, err := control.NewService(control.Options{
|
||||
Store: persistence, DefaultQueueTTL: configured.Queue.DefaultTTL, DefaultQueueTTLSet: true,
|
||||
Store: persistence, DefaultQueueTTL: configured.Queue.DefaultTTL, DefaultQueueTTLSet: true, TakeoverTTL: configured.Protocol.TakeoverTTL,
|
||||
WakeClient: registry.Wake,
|
||||
Limits: agentproto.Limits{MaxEnvelopeBytes: configured.Protocol.MaxAgentEnvelopeBytes, MaxExecutionSpecBytes: configured.Protocol.MaxExecutionSpecBytes, MaxRawChunkBytes: configured.Protocol.MaxRawChunkBytes, MaxScriptBytes: configured.Protocol.MaxScriptBytes, MaxDetailBytes: configured.Storage.ProtocolDetailMaxBytes},
|
||||
})
|
||||
|
||||
+251
-1
@@ -16,6 +16,7 @@ import (
|
||||
"time"
|
||||
|
||||
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
|
||||
"github.com/rvbox/rvbox/internal/agentproto"
|
||||
"github.com/rvbox/rvbox/internal/domain"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
@@ -55,6 +56,16 @@ func run(args []string, output, diagnostics io.Writer) error {
|
||||
return stat(ctx, client, args[1:], output)
|
||||
case "run":
|
||||
return runCommand(ctx, client, args[1:], output, diagnostics)
|
||||
case "append":
|
||||
return appendStdin(ctx, client, args[1:], output)
|
||||
case "close-stdin":
|
||||
return closeStdin(ctx, client, args[1:], output)
|
||||
case "kill":
|
||||
return killCommand(ctx, client, args[1:], output)
|
||||
case "storage":
|
||||
return storage(ctx, client, args[1:], output)
|
||||
case "client":
|
||||
return clientCommand(ctx, client, args[1:], output)
|
||||
default:
|
||||
return fmt.Errorf("unknown command %q", args[0])
|
||||
}
|
||||
@@ -224,14 +235,253 @@ func runCommand(ctx context.Context, client rvboxv1.ControlClient, args []string
|
||||
return printRunResponse(ctx, client, request, output, *background)
|
||||
}
|
||||
|
||||
func printRunResponse(ctx context.Context, client rvboxv1.ControlClient, request *rvboxv1.RunCommandRequest, output io.Writer, _ bool) error {
|
||||
func printRunResponse(ctx context.Context, client rvboxv1.ControlClient, request *rvboxv1.RunCommandRequest, output io.Writer, background bool) error {
|
||||
response, err := client.RunCommand(ctx, request)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintf(output, "%s %s\n", response.GetIssueUuid(), response.GetLifecycle())
|
||||
if background {
|
||||
return nil
|
||||
}
|
||||
follow, err := client.FollowCommand(ctx, &rvboxv1.FollowCommandRequest{ClientId: request.GetTargetClientId(), IssueUuid: response.GetIssueUuid(), IncludeExisting: true})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for {
|
||||
item, err := follow.Recv()
|
||||
if errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
event := item.GetEvent()
|
||||
if event == nil {
|
||||
continue
|
||||
}
|
||||
switch payload := event.Payload.(type) {
|
||||
case *rvboxv1.CommandEvent_Output:
|
||||
data, decodeErr := agentproto.DecodeOutputChunk(payload.Output, 64<<10)
|
||||
if decodeErr != nil {
|
||||
return decodeErr
|
||||
}
|
||||
if _, writeErr := output.Write(data); writeErr != nil {
|
||||
return writeErr
|
||||
}
|
||||
case *rvboxv1.CommandEvent_Lifecycle:
|
||||
fmt.Fprintf(output, "\n[%s] %s\n", payload.Lifecycle.GetLifecycle(), payload.Lifecycle.GetDetail())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func appendStdin(ctx context.Context, client rvboxv1.ControlClient, args []string, output io.Writer) error {
|
||||
flags := flag.NewFlagSet("append", flag.ContinueOnError)
|
||||
flags.SetOutput(io.Discard)
|
||||
filePath := flags.String("file", "", "read stdin data from a file")
|
||||
raw := flags.Bool("raw", false, "do not append a newline")
|
||||
requestID := flags.String("request-id", "", "canonical UUIDv7")
|
||||
if err := flags.Parse(args); err != nil {
|
||||
return err
|
||||
}
|
||||
positionals := flags.Args()
|
||||
if len(positionals) < 2 || len(positionals) > 3 || (*filePath != "" && len(positionals) == 3) {
|
||||
return errors.New("append requires CLIENT UUID TEXT, or --file PATH CLIENT UUID")
|
||||
}
|
||||
data := []byte{}
|
||||
var err error
|
||||
if *filePath != "" {
|
||||
data, err = os.ReadFile(*filePath)
|
||||
} else {
|
||||
data = []byte(positionals[2])
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if *requestID == "" {
|
||||
*requestID, err = generatedRequestID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
issue := positionals[1]
|
||||
if _, err := domain.ParseUUIDv7(issue); err != nil {
|
||||
return err
|
||||
}
|
||||
response, err := client.AppendStdin(ctx, &rvboxv1.AppendStdinRequest{ClientId: positionals[0], IssueUuid: issue, Data: data, AppendNewline: !*raw, RequestId: *requestID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintf(output, "%d\n", response.GetWriteSeq())
|
||||
return nil
|
||||
}
|
||||
|
||||
func closeStdin(ctx context.Context, client rvboxv1.ControlClient, args []string, output io.Writer) error {
|
||||
flags := flag.NewFlagSet("close-stdin", flag.ContinueOnError)
|
||||
flags.SetOutput(io.Discard)
|
||||
requestID := flags.String("request-id", "", "canonical UUIDv7")
|
||||
if err := flags.Parse(args); err != nil {
|
||||
return err
|
||||
}
|
||||
if len(flags.Args()) != 2 {
|
||||
return errors.New("close-stdin requires CLIENT UUID")
|
||||
}
|
||||
if _, err := domain.ParseUUIDv7(flags.Args()[1]); err != nil {
|
||||
return err
|
||||
}
|
||||
if *requestID == "" {
|
||||
var err error
|
||||
*requestID, err = generatedRequestID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
response, err := client.CloseStdin(ctx, &rvboxv1.CloseStdinRequest{ClientId: flags.Args()[0], IssueUuid: flags.Args()[1], RequestId: *requestID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintf(output, "%d\n", response.GetWriteSeq())
|
||||
return nil
|
||||
}
|
||||
|
||||
func killCommand(ctx context.Context, client rvboxv1.ControlClient, args []string, output io.Writer) error {
|
||||
flags := flag.NewFlagSet("kill", flag.ContinueOnError)
|
||||
flags.SetOutput(io.Discard)
|
||||
requestID := flags.String("request-id", "", "canonical UUIDv7")
|
||||
if err := flags.Parse(args); err != nil {
|
||||
return err
|
||||
}
|
||||
positionals := flags.Args()
|
||||
signal := rvboxv1.SignalKind_SIGNAL_TERM
|
||||
if len(positionals) == 3 {
|
||||
parsed, err := parseSignal(positionals[0])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
signal = parsed
|
||||
positionals = positionals[1:]
|
||||
}
|
||||
if len(positionals) != 2 {
|
||||
return errors.New("kill requires [SIGNAL] CLIENT UUID")
|
||||
}
|
||||
if _, err := domain.ParseUUIDv7(positionals[1]); err != nil {
|
||||
return err
|
||||
}
|
||||
if *requestID == "" {
|
||||
var err error
|
||||
*requestID, err = generatedRequestID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
response, err := client.SignalCommand(ctx, &rvboxv1.ControlSignalCommandRequest{ClientId: positionals[0], IssueUuid: positionals[1], Signal: signal, RequestId: *requestID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintf(output, "%d\n", response.GetCommandRevision())
|
||||
return nil
|
||||
}
|
||||
|
||||
func storage(ctx context.Context, client rvboxv1.ControlClient, args []string, output io.Writer) error {
|
||||
if len(args) == 0 {
|
||||
return errors.New("storage requires incidents, repair, or acknowledge")
|
||||
}
|
||||
switch args[0] {
|
||||
case "incidents":
|
||||
response, err := client.ListStorageIncidents(ctx, &rvboxv1.ListStorageIncidentsRequest{IncludeResolved: contains(args[1:], "--all")})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, incident := range response.GetIncidents() {
|
||||
fmt.Fprintf(output, "%s state=%s scope=%s summary=%s\n", incident.GetIncidentId(), incident.GetState(), incident.GetScope(), incident.GetSummary())
|
||||
}
|
||||
return nil
|
||||
case "repair", "acknowledge":
|
||||
if len(args) < 2 {
|
||||
return errors.New("storage mutation requires INCIDENT_ID")
|
||||
}
|
||||
requestID, err := generatedRequestID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if args[0] == "repair" {
|
||||
response, callErr := client.RepairStorageIncident(ctx, &rvboxv1.RepairStorageIncidentRequest{IncidentId: args[1], RequestId: requestID})
|
||||
if callErr != nil {
|
||||
return callErr
|
||||
}
|
||||
fmt.Fprintf(output, "%s\n", response.GetIncident().GetState())
|
||||
return nil
|
||||
}
|
||||
note := strings.Join(args[2:], " ")
|
||||
if note == "" {
|
||||
return errors.New("storage acknowledge requires a note")
|
||||
}
|
||||
response, callErr := client.AcknowledgeStorageIncident(ctx, &rvboxv1.AcknowledgeStorageIncidentRequest{IncidentId: args[1], RequestId: requestID, Note: note})
|
||||
if callErr != nil {
|
||||
return callErr
|
||||
}
|
||||
fmt.Fprintf(output, "%s\n", response.GetIncident().GetState())
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("unknown storage command %q", args[0])
|
||||
}
|
||||
}
|
||||
|
||||
func clientCommand(ctx context.Context, client rvboxv1.ControlClient, args []string, output io.Writer) error {
|
||||
if len(args) != 3 || args[0] != "takeover" {
|
||||
return errors.New("client takeover requires CLIENT INSTANCE_ID")
|
||||
}
|
||||
if _, err := domain.ParseUUIDv7(args[2]); err != nil {
|
||||
return err
|
||||
}
|
||||
requestID, err := generatedRequestID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
response, err := client.AuthorizeClientTakeover(ctx, &rvboxv1.AuthorizeClientTakeoverRequest{ClientId: args[1], ClientInstanceId: args[2], RequestId: requestID})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
fmt.Fprintln(output, response.GetExpiresAt().AsTime().UTC().Format(time.RFC3339Nano))
|
||||
return nil
|
||||
}
|
||||
|
||||
func generatedRequestID() (string, error) {
|
||||
value, err := domain.NewUUIDv7()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return value.String(), nil
|
||||
}
|
||||
|
||||
func parseSignal(value string) (rvboxv1.SignalKind, error) {
|
||||
value = strings.ToUpper(strings.TrimPrefix(value, "SIG"))
|
||||
switch value {
|
||||
case "HUP":
|
||||
return rvboxv1.SignalKind_SIGNAL_HUP, nil
|
||||
case "INT":
|
||||
return rvboxv1.SignalKind_SIGNAL_INT, nil
|
||||
case "TERM":
|
||||
return rvboxv1.SignalKind_SIGNAL_TERM, nil
|
||||
case "KILL":
|
||||
return rvboxv1.SignalKind_SIGNAL_KILL, nil
|
||||
case "USR1":
|
||||
return rvboxv1.SignalKind_SIGNAL_USR1, nil
|
||||
case "USR2":
|
||||
return rvboxv1.SignalKind_SIGNAL_USR2, nil
|
||||
default:
|
||||
return rvboxv1.SignalKind_SIGNAL_KIND_UNSPECIFIED, fmt.Errorf("unsupported signal %q", value)
|
||||
}
|
||||
}
|
||||
|
||||
func contains(values []string, target string) bool {
|
||||
for _, value := range values {
|
||||
if value == target {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func queueDuration(value time.Duration) *durationpb.Duration {
|
||||
if value < 0 {
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
|
||||
)
|
||||
|
||||
func TestGlobalSocketAndCLIValueParsing_HP_CTL_11(t *testing.T) {
|
||||
socket, args, err := globalSocket([]string{"--socket", filepath.Join(t.TempDir(), "rvbox.sock"), "run", "client", "echo"})
|
||||
if err != nil || args[0] != "run" || socket == "" {
|
||||
t.Fatalf("global socket = %q %#v, %v", socket, args, err)
|
||||
}
|
||||
if _, _, err := globalSocket([]string{"--socket", "relative.sock", "stat"}); err == nil {
|
||||
t.Fatal("relative control socket accepted")
|
||||
}
|
||||
for input, want := range map[string]rvboxv1.ShellType{"sh": rvboxv1.ShellType_SHELL_SH, "bash": rvboxv1.ShellType_SHELL_BASH, "cmd": rvboxv1.ShellType_SHELL_CMD, "pwsh": rvboxv1.ShellType_SHELL_POWERSHELL, "": rvboxv1.ShellType_SHELL_TYPE_UNSPECIFIED} {
|
||||
got, err := parseShell(input)
|
||||
if err != nil || got != want {
|
||||
t.Errorf("parseShell(%q) = %v, %v", input, got, err)
|
||||
}
|
||||
}
|
||||
if _, err := parseProfile("unknown"); err == nil {
|
||||
t.Fatal("unknown profile accepted")
|
||||
}
|
||||
if got := queueDuration(-time.Nanosecond); got != nil {
|
||||
t.Fatal("negative queue duration was encoded")
|
||||
}
|
||||
if got := queueDuration(0); got == nil || got.AsDuration() != 0 {
|
||||
t.Fatalf("zero queue duration = %v", got)
|
||||
}
|
||||
}
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"sort"
|
||||
"time"
|
||||
|
||||
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
|
||||
@@ -41,6 +42,7 @@ type Options struct {
|
||||
// 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
|
||||
@@ -51,6 +53,7 @@ type Service struct {
|
||||
rvboxv1.UnimplementedControlServer
|
||||
store *store.Store
|
||||
defaultQueueTTL time.Duration
|
||||
takeoverTTL time.Duration
|
||||
limits agentproto.Limits
|
||||
now func() time.Time
|
||||
cursors *domain.CursorCodec
|
||||
@@ -67,6 +70,12 @@ func NewService(options Options) (*Service, error) {
|
||||
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()
|
||||
}
|
||||
@@ -84,7 +93,7 @@ func NewService(options Options) (*Service, error) {
|
||||
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
|
||||
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) {
|
||||
@@ -272,6 +281,428 @@ func (service *Service) RunCommand(ctx context.Context, request *rvboxv1.RunComm
|
||||
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()
|
||||
@@ -403,6 +834,10 @@ func mapStoreError(err error) error {
|
||||
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:
|
||||
|
||||
@@ -15,10 +15,12 @@ import (
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/grpc/test/bufconn"
|
||||
"google.golang.org/protobuf/proto"
|
||||
"google.golang.org/protobuf/types/known/durationpb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
)
|
||||
|
||||
func TestControlListAndGetViews_HP_CONTROL_01(t *testing.T) {
|
||||
@@ -147,6 +149,174 @@ func TestControlGRPCRoundTrip_HP_CONTROL_06(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFollowCommandStreamsDurableEvents_HP_CONTROL_10(t *testing.T) {
|
||||
service, persistence := newTestService(t)
|
||||
defer persistence.Close()
|
||||
registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_CMD, 6)
|
||||
issue := fixedIssue(0xa6)
|
||||
now := time.Now().UTC()
|
||||
specBytes, err := proto.Marshal(&rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_CMD, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo"}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := persistence.QueueCommand(context.Background(), store.QueueCommandInput{IssueUUID: issue, ClientID: "win-a", IssueTime: now, ReceiptTime: now, ImmutableSHA256: sha256.Sum256([]byte("follow")), ExecutionSpec: specBytes}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := persistence.ClaimNextDispatch(context.Background(), "win-a", 1, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := persistence.RecordCommandAcceptance(context.Background(), issue, "win-a", 1, 1, true, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for index, lifecycle := range []rvboxv1.CommandLifecycle{rvboxv1.CommandLifecycle_COMMAND_RUNNING, rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED} {
|
||||
event := &rvboxv1.CommandEvent{IssueUuid: issue.String(), EventSeq: uint64(index + 1), ObservedAt: timestamppb.New(now.Add(time.Duration(index) * time.Millisecond)), Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: &rvboxv1.LifecycleChange{Lifecycle: lifecycle, CommandRevision: 1}}}
|
||||
digest, digestErr := agentproto.CommandEventDigest(event)
|
||||
if digestErr != nil {
|
||||
t.Fatal(digestErr)
|
||||
}
|
||||
event.ImmutableEventSha256 = digest[:]
|
||||
payload, marshalErr := proto.MarshalOptions{Deterministic: true}.Marshal(event)
|
||||
if marshalErr != nil {
|
||||
t.Fatal(marshalErr)
|
||||
}
|
||||
if _, appendErr := persistence.AppendCommandEvent(context.Background(), store.EventAppend{IssueUUID: [16]byte(issue), ClientID: "win-a", SessionGeneration: 1, EventSeq: uint64(index + 1), ObservedUnixNano: event.GetObservedAt().AsTime().UnixNano(), ReceiptUnixNano: now.Add(time.Duration(index) * time.Millisecond).UnixNano(), EventType: 4, Compression: 1, RawLength: uint64(len(payload)), Payload: payload, ImmutableSHA256: digest, Lifecycle: &lifecycle, LifecycleRevision: 1}); appendErr != nil {
|
||||
t.Fatal(appendErr)
|
||||
}
|
||||
}
|
||||
stream := &testFollowStream{ctx: context.Background()}
|
||||
if err := service.FollowCommand(&rvboxv1.FollowCommandRequest{ClientId: "win-a", IssueUuid: issue.String(), IncludeExisting: true}, stream); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(stream.responses) != 2 || stream.responses[0].GetEvent().GetEventSeq() != 1 || stream.responses[1].GetEvent().GetLifecycle().GetLifecycle() != rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED {
|
||||
t.Fatalf("follow responses = %#v", stream.responses)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetOutputSlicesWithAuthenticatedCursor_HP_CONTROL_12(t *testing.T) {
|
||||
service, persistence := newTestService(t)
|
||||
defer persistence.Close()
|
||||
registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_CMD, 7)
|
||||
issue := fixedIssue(0xa7)
|
||||
now := time.Now().UTC()
|
||||
specBytes, _ := proto.Marshal(&rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_CMD, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo"}})
|
||||
if _, err := persistence.QueueCommand(context.Background(), store.QueueCommandInput{IssueUUID: issue, ClientID: "win-a", IssueTime: now, ReceiptTime: now, ImmutableSHA256: sha256.Sum256([]byte("output")), ExecutionSpec: specBytes}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := persistence.ClaimNextDispatch(context.Background(), "win-a", 1, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := persistence.RecordCommandAcceptance(context.Background(), issue, "win-a", 1, 1, true, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
event := &rvboxv1.CommandEvent{IssueUuid: issue.String(), EventSeq: 1, ObservedAt: timestamppb.New(now), Payload: &rvboxv1.CommandEvent_Output{Output: &rvboxv1.OutputChunk{Stream: rvboxv1.StreamKind_STREAM_STDOUT, Compression: rvboxv1.Compression_COMPRESSION_NONE, Data: []byte("abcdef"), UncompressedSize: 6, CompressedSize: 6}}}
|
||||
digest, err := agentproto.CommandEventDigest(event)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
event.ImmutableEventSha256 = digest[:]
|
||||
payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(event)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := persistence.AppendCommandEvent(context.Background(), store.EventAppend{IssueUUID: [16]byte(issue), ClientID: "win-a", SessionGeneration: 1, EventSeq: 1, ObservedUnixNano: now.UnixNano(), ReceiptUnixNano: now.UnixNano(), EventType: 5, Stream: 1, Compression: 1, RawLength: uint64(len(payload)), Payload: payload, ImmutableSHA256: digest, Output: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
first, err := service.GetOutput(context.Background(), &rvboxv1.GetOutputRequest{ClientId: "win-a", IssueUuid: issue.String(), Streams: []rvboxv1.StreamKind{rvboxv1.StreamKind_STREAM_STDOUT}, MaxBytes: 3})
|
||||
if err != nil || len(first.GetOutput()) != 1 || string(first.GetOutput()[0].GetData()) != "abc" || first.GetNextPageToken() == "" {
|
||||
t.Fatalf("first output page = %#v, %v", first, err)
|
||||
}
|
||||
second, err := service.GetOutput(context.Background(), &rvboxv1.GetOutputRequest{ClientId: "win-a", IssueUuid: issue.String(), Streams: []rvboxv1.StreamKind{rvboxv1.StreamKind_STREAM_STDOUT}, MaxBytes: 3, PageToken: first.GetNextPageToken()})
|
||||
if err != nil || len(second.GetOutput()) != 1 || string(second.GetOutput()[0].GetData()) != "def" || !second.GetOutput()[0].GetEndOfEvent() {
|
||||
t.Fatalf("second output page = %#v, %v", second, err)
|
||||
}
|
||||
if _, err := service.GetOutput(context.Background(), &rvboxv1.GetOutputRequest{ClientId: "win-a", IssueUuid: issue.String(), Streams: []rvboxv1.StreamKind{rvboxv1.StreamKind_STREAM_STDERR}, MaxBytes: 3, PageToken: first.GetNextPageToken()}); status.Code(err) != codes.InvalidArgument {
|
||||
t.Fatalf("changed output filter code = %v", status.Code(err))
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlMutationIdempotencyAndQueuedCancellation_BH_CONTROL_13(t *testing.T) {
|
||||
service, persistence := newTestService(t)
|
||||
defer persistence.Close()
|
||||
registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_CMD, 8)
|
||||
issue := fixedIssue(0xa8)
|
||||
if _, err := service.RunCommand(context.Background(), &rvboxv1.RunCommandRequest{TargetClientId: "win-a", RequestId: issue.String(), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_CMD, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "wait"}}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stdinID := fixedIssue(0xa9)
|
||||
appendRequest := &rvboxv1.AppendStdinRequest{ClientId: "win-a", IssueUuid: issue.String(), RequestId: stdinID.String(), Data: []byte("hello"), AppendNewline: true}
|
||||
first, err := service.AppendStdin(context.Background(), appendRequest)
|
||||
if err != nil || first.GetWriteSeq() != 1 {
|
||||
t.Fatalf("append stdin = %#v, %v", first, err)
|
||||
}
|
||||
second, err := service.AppendStdin(context.Background(), proto.Clone(appendRequest).(*rvboxv1.AppendStdinRequest))
|
||||
if err != nil || second.GetWriteSeq() != first.GetWriteSeq() {
|
||||
t.Fatalf("append replay = %#v, %v", second, err)
|
||||
}
|
||||
conflict := proto.Clone(appendRequest).(*rvboxv1.AppendStdinRequest)
|
||||
conflict.Data = []byte("different")
|
||||
if _, err := service.AppendStdin(context.Background(), conflict); status.Code(err) != codes.AlreadyExists {
|
||||
t.Fatalf("append conflict code = %v", status.Code(err))
|
||||
}
|
||||
closeID := fixedIssue(0xaa)
|
||||
closed, err := service.CloseStdin(context.Background(), &rvboxv1.CloseStdinRequest{ClientId: "win-a", IssueUuid: issue.String(), RequestId: closeID.String()})
|
||||
if err != nil || closed.GetWriteSeq() != 2 {
|
||||
t.Fatalf("close stdin = %#v, %v", closed, err)
|
||||
}
|
||||
signalID := fixedIssue(0xab)
|
||||
cancelled, err := service.SignalCommand(context.Background(), &rvboxv1.ControlSignalCommandRequest{ClientId: "win-a", IssueUuid: issue.String(), RequestId: signalID.String(), Signal: rvboxv1.SignalKind_SIGNAL_TERM})
|
||||
if err != nil || cancelled.GetCommandRevision() != 2 {
|
||||
t.Fatalf("queued cancellation = %#v, %v", cancelled, err)
|
||||
}
|
||||
replay, err := service.SignalCommand(context.Background(), &rvboxv1.ControlSignalCommandRequest{ClientId: "win-a", IssueUuid: issue.String(), RequestId: signalID.String(), Signal: rvboxv1.SignalKind_SIGNAL_TERM})
|
||||
if err != nil || replay.GetCommandRevision() != cancelled.GetCommandRevision() {
|
||||
t.Fatalf("cancel replay = %#v, %v", replay, err)
|
||||
}
|
||||
command, err := service.GetCommand(context.Background(), &rvboxv1.GetCommandRequest{ClientId: "win-a", IssueUuid: issue.String()})
|
||||
if err != nil || command.GetCommand().GetLifecycle() != rvboxv1.CommandLifecycle_COMMAND_CANCELLED {
|
||||
t.Fatalf("cancelled command = %#v, %v", command, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStorageIncidentControlLifecycle_HP_CONTROL_14(t *testing.T) {
|
||||
service, persistence := newTestService(t)
|
||||
defer persistence.Close()
|
||||
incidentID := fixedIssue(0xac)
|
||||
if _, err := persistence.RecordIncident(context.Background(), store.IncidentInput{IncidentUUID: [16]byte(incidentID), DetectedAt: time.Now().UTC(), Kind: store.IncidentChecksumMismatch, Scope: store.IncidentScopeGlobal, ScopeKey: "segments", Summary: "checksum mismatch", Evidence: []byte("evidence"), AutomaticallyRepairable: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
listed, err := service.ListStorageIncidents(context.Background(), &rvboxv1.ListStorageIncidentsRequest{PageSize: 1})
|
||||
if err != nil || len(listed.GetIncidents()) != 1 || listed.GetIncidents()[0].GetState() != rvboxv1.StorageIncidentState_STORAGE_INCIDENT_STATE_OPEN {
|
||||
t.Fatalf("incident list = %#v, %v", listed, err)
|
||||
}
|
||||
repaired, err := service.RepairStorageIncident(context.Background(), &rvboxv1.RepairStorageIncidentRequest{IncidentId: incidentID.String(), RequestId: fixedIssue(0xad).String()})
|
||||
if err != nil || repaired.GetIncident().GetState() != rvboxv1.StorageIncidentState_STORAGE_INCIDENT_STATE_REPAIRED {
|
||||
t.Fatalf("incident repair = %#v, %v", repaired, err)
|
||||
}
|
||||
if repeat, err := service.RepairStorageIncident(context.Background(), &rvboxv1.RepairStorageIncidentRequest{IncidentId: incidentID.String(), RequestId: fixedIssue(0xae).String()}); err != nil || repeat.GetIncident().GetState() != rvboxv1.StorageIncidentState_STORAGE_INCIDENT_STATE_REPAIRED {
|
||||
t.Fatalf("second repair = %#v, %v", repeat, err)
|
||||
}
|
||||
resolved, err := service.ListStorageIncidents(context.Background(), &rvboxv1.ListStorageIncidentsRequest{IncludeResolved: true})
|
||||
if err != nil || len(resolved.GetIncidents()) != 1 || resolved.GetIncidents()[0].GetResolvedAt() == nil {
|
||||
t.Fatalf("resolved incident list = %#v, %v", resolved, err)
|
||||
}
|
||||
}
|
||||
|
||||
type testFollowStream struct {
|
||||
ctx context.Context
|
||||
responses []*rvboxv1.FollowCommandResponse
|
||||
}
|
||||
|
||||
func (stream *testFollowStream) Send(response *rvboxv1.FollowCommandResponse) error {
|
||||
stream.responses = append(stream.responses, response)
|
||||
return nil
|
||||
}
|
||||
func (stream *testFollowStream) SetHeader(metadata.MD) error { return nil }
|
||||
func (stream *testFollowStream) SendHeader(metadata.MD) error { return nil }
|
||||
func (stream *testFollowStream) SetTrailer(metadata.MD) {}
|
||||
func (stream *testFollowStream) Context() context.Context { return stream.ctx }
|
||||
func (stream *testFollowStream) SendMsg(any) error { return nil }
|
||||
func (stream *testFollowStream) RecvMsg(any) error { return nil }
|
||||
|
||||
func newTestService(t *testing.T) (*Service, *store.Store) {
|
||||
t.Helper()
|
||||
persistence, err := store.Open(context.Background(), store.Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second})
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/klauspost/compress/zstd"
|
||||
"github.com/rvbox/rvbox/internal/domain"
|
||||
)
|
||||
|
||||
// EventView is a validated persisted command event. Payload is the stored
|
||||
// event bytes (normally deterministic protobuf bytes); callers that need the
|
||||
// uncompressed form should use DecodeEventPayload.
|
||||
type EventView struct {
|
||||
IssueUUID domain.UUID
|
||||
EventSeq uint64
|
||||
ObservedUnixNano int64
|
||||
ReceiptUnixNano int64
|
||||
EventType uint16
|
||||
Stream uint16
|
||||
Compression uint16
|
||||
RawLength uint64
|
||||
Payload []byte
|
||||
ImmutableSHA256 [32]byte
|
||||
}
|
||||
|
||||
// ReadCommandEvents returns at most limit events after afterEventSeq in strict
|
||||
// sequence order. Segment references are revalidated at read time; a malformed
|
||||
// or replaced segment is an error rather than silently returning partial data.
|
||||
func (store *Store) ReadCommandEvents(ctx context.Context, issue domain.UUID, afterEventSeq uint64, limit uint32) ([]EventView, error) {
|
||||
if issue == (domain.UUID{}) || limit == 0 || limit > 1000 {
|
||||
return nil, errors.New("invalid command event read")
|
||||
}
|
||||
database, err := store.openDatabase()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rows, err := database.QueryContext(ctx, `SELECT
|
||||
ce.event_seq, ce.observed_at, ce.server_receipt_time, ce.event_type,
|
||||
ce.compression, ce.raw_bytes, ce.payload, ce.segment_ordinal,
|
||||
ce.segment_record_offset, ce.segment_record_length, os.path,
|
||||
ce.immutable_sha256
|
||||
FROM command_events ce LEFT JOIN output_segments os
|
||||
ON os.issue_uuid = ce.issue_uuid AND os.ordinal = ce.segment_ordinal
|
||||
WHERE ce.issue_uuid = ? AND ce.event_seq > ? ORDER BY ce.event_seq LIMIT ?`, issue[:], afterEventSeq, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := make([]EventView, 0, limit)
|
||||
for rows.Next() {
|
||||
var event EventView
|
||||
var inline []byte
|
||||
var ordinal, offset, length sql.NullInt64
|
||||
var path sql.NullString
|
||||
var digest []byte
|
||||
if err := rows.Scan(&event.EventSeq, &event.ObservedUnixNano, &event.ReceiptUnixNano, &event.EventType, &event.Compression, &event.RawLength, &inline, &ordinal, &offset, &length, &path, &digest); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(digest) != 32 || event.EventSeq == 0 || event.EventType == 0 {
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
copy(event.IssueUUID[:], issue[:])
|
||||
copy(event.ImmutableSHA256[:], digest)
|
||||
switch {
|
||||
case inline != nil && !ordinal.Valid && !offset.Valid && !length.Valid && !path.Valid:
|
||||
event.Payload = append([]byte(nil), inline...)
|
||||
case inline == nil && ordinal.Valid && offset.Valid && length.Valid && path.Valid:
|
||||
if ordinal.Int64 < 0 || offset.Int64 < 0 || length.Int64 <= 0 || ordinal.Int64 > int64(^uint32(0)) || uint64(length.Int64) > uint64(^uint(0)>>1) {
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
payload, stream, readErr := store.readEventSegment(path.String, issue, event.EventSeq, uint32(ordinal.Int64), uint64(offset.Int64), uint64(length.Int64))
|
||||
if readErr != nil {
|
||||
return nil, readErr
|
||||
}
|
||||
event.Payload = payload
|
||||
event.Stream = stream
|
||||
default:
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
if uint64(len(event.Payload)) > DefaultSegmentLimit || event.Compression < 1 || event.Compression > 2 || event.RawLength > DefaultSegmentLimit {
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
result = append(result, event)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (store *Store) readEventSegment(name string, issue domain.UUID, sequence uint64, ordinal uint32, offset, length uint64) ([]byte, uint16, error) {
|
||||
if filepath.Base(name) != name {
|
||||
return nil, 0, ErrUnsafeSegmentReference
|
||||
}
|
||||
owner, parsedOrdinal, err := parseSegmentName(name)
|
||||
if err != nil || owner != [16]byte(issue) || parsedOrdinal != ordinal {
|
||||
return nil, 0, ErrUnsafeSegmentReference
|
||||
}
|
||||
path := filepath.Join(store.segmentDirectory(), name)
|
||||
info, err := os.Lstat(path)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if !info.Mode().IsRegular() || info.Mode()&os.ModeSymlink != 0 || info.Mode().Perm()&0o077 != 0 {
|
||||
return nil, 0, ErrUnsafeSegmentReference
|
||||
}
|
||||
if length > uint64(^uint(0)>>1) || offset > uint64(^uint(0)>>1) || offset+length < offset || offset+length > uint64(info.Size()) {
|
||||
return nil, 0, ErrCommittedRangeMissing
|
||||
}
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer file.Close()
|
||||
reader := io.NewSectionReader(file, int64(offset), int64(length))
|
||||
record, encodedLength, err := DecodeSegmentRecord(reader, DefaultSegmentLimits())
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("decode command event segment: %w", err)
|
||||
}
|
||||
if encodedLength != length || record.OwnerUUID != [16]byte(issue) || record.Sequence != sequence || record.Kind != PayloadKindCommandEvent {
|
||||
return nil, 0, ErrInvalidSegmentRecord
|
||||
}
|
||||
return append([]byte(nil), record.Payload...), record.Stream, nil
|
||||
}
|
||||
|
||||
// DecodeEventPayload validates and boundedly decompresses a stored event
|
||||
// payload. It is intentionally generic so output/follow handlers share the
|
||||
// same decompression ceiling.
|
||||
func DecodeEventPayload(event EventView) ([]byte, error) {
|
||||
if event.Compression == 1 {
|
||||
if event.RawLength != uint64(len(event.Payload)) {
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
return append([]byte(nil), event.Payload...), nil
|
||||
}
|
||||
if event.Compression != 2 || event.RawLength > DefaultSegmentLimit {
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
decoder, err := zstd.NewReader(bytes.NewReader(event.Payload), zstd.WithDecoderConcurrency(1), zstd.WithDecoderMaxMemory(DefaultSegmentLimit+1))
|
||||
if err != nil {
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
defer decoder.Close()
|
||||
decoded, err := io.ReadAll(io.LimitReader(decoder, int64(event.RawLength)+1))
|
||||
if err != nil || uint64(len(decoded)) != event.RawLength {
|
||||
return nil, ErrInvalidSegmentRecord
|
||||
}
|
||||
return decoded, nil
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
|
||||
"github.com/rvbox/rvbox/internal/domain"
|
||||
)
|
||||
|
||||
func TestReadCommandEventsValidatesSegmentAndPreservesOrder_HP_EVENT_04(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
opened, err := Open(ctx, Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer opened.Close()
|
||||
if _, err := opened.RegisterClientSession(ctx, ClientRegistration{ClientID: "win-client", Platform: 3, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\`, SupportedShells: []byte{1}, ClientInstanceID: [16]byte{1}, SessionID: [16]byte{2}, ConnectedAt: time.Now()}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-0000000000b1")
|
||||
now := time.Now().UTC()
|
||||
if _, err := opened.QueueCommand(ctx, QueueCommandInput{IssueUUID: issue, ClientID: "win-client", IssueTime: now, ReceiptTime: now, ImmutableSHA256: sha256.Sum256([]byte("request")), ExecutionSpec: []byte("spec")}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := opened.ClaimNextDispatch(ctx, "win-client", 1, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := opened.RecordCommandAcceptance(ctx, issue, "win-client", 1, 1, true, now); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for sequence, lifecycle := range []rvboxv1.CommandLifecycle{rvboxv1.CommandLifecycle_COMMAND_RUNNING, rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED} {
|
||||
event := []byte{byte(sequence + 1), 'e', 'v', 't'}
|
||||
if _, err := opened.AppendCommandEvent(ctx, EventAppend{IssueUUID: [16]byte(issue), ClientID: "win-client", SessionGeneration: 1, EventSeq: uint64(sequence + 1), ObservedUnixNano: now.UnixNano(), ReceiptUnixNano: now.Add(time.Duration(sequence) * time.Millisecond).UnixNano(), EventType: 4, Compression: 1, RawLength: uint64(len(event)), Payload: event, ImmutableSHA256: sha256.Sum256(event), Lifecycle: &lifecycle, LifecycleRevision: 1}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
events, err := opened.ReadCommandEvents(ctx, issue, 0, 10)
|
||||
if err != nil || len(events) != 2 || events[0].EventSeq != 1 || events[1].EventSeq != 2 || string(events[1].Payload) != "\x02evt" {
|
||||
t.Fatalf("read events = %#v, %v", events, err)
|
||||
}
|
||||
decoded, err := DecodeEventPayload(events[0])
|
||||
if err != nil || string(decoded) != "\x01evt" {
|
||||
t.Fatalf("decoded event payload = %q, %v", decoded, err)
|
||||
}
|
||||
if _, err := opened.ReadCommandEvents(ctx, issue, 0, 0); err == nil {
|
||||
t.Fatal("zero event limit accepted")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
type IncidentView struct {
|
||||
IncidentUUID [16]byte
|
||||
DetectedAt time.Time
|
||||
ResolvedAt *time.Time
|
||||
State IncidentState
|
||||
Kind IncidentKind
|
||||
Scope IncidentScope
|
||||
ClientID string
|
||||
IssueUUID *[16]byte
|
||||
Summary string
|
||||
DataLoss bool
|
||||
AutomaticallyRepairable bool
|
||||
}
|
||||
|
||||
type IncidentPage struct {
|
||||
IncludeResolved bool
|
||||
Limit uint32
|
||||
SnapshotBoundary int64
|
||||
AfterDetectedAt int64
|
||||
AfterUUID [16]byte
|
||||
HasAfter bool
|
||||
}
|
||||
|
||||
func (store *Store) ListIncidentViews(ctx context.Context, page IncidentPage) ([]IncidentView, bool, error) {
|
||||
if page.Limit == 0 || page.Limit > 1000 || page.SnapshotBoundary <= 0 {
|
||||
return nil, false, errors.New("invalid incident page")
|
||||
}
|
||||
database, err := store.openDatabase()
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
query := `SELECT incident_uuid, detected_at, resolved_at, state, kind, scope, client_id, issue_uuid, summary, data_loss, automatically_repairable FROM storage_incidents WHERE detected_at <= ?`
|
||||
args := []any{page.SnapshotBoundary}
|
||||
if !page.IncludeResolved {
|
||||
query += ` AND state = 1`
|
||||
}
|
||||
if page.HasAfter {
|
||||
query += ` AND (detected_at < ? OR (detected_at = ? AND incident_uuid < ?))`
|
||||
args = append(args, page.AfterDetectedAt, page.AfterDetectedAt, page.AfterUUID[:])
|
||||
}
|
||||
query += ` ORDER BY detected_at DESC, incident_uuid DESC LIMIT ?`
|
||||
args = append(args, page.Limit+1)
|
||||
rows, err := database.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
defer rows.Close()
|
||||
result := make([]IncidentView, 0, page.Limit)
|
||||
for rows.Next() {
|
||||
view, scanErr := scanIncidentView(rows)
|
||||
if scanErr != nil {
|
||||
return nil, false, scanErr
|
||||
}
|
||||
if uint32(len(result)) < page.Limit {
|
||||
result = append(result, view)
|
||||
} else {
|
||||
return result, true, rows.Err()
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
return result, false, nil
|
||||
}
|
||||
|
||||
func (store *Store) GetIncidentView(ctx context.Context, incidentUUID [16]byte) (IncidentView, error) {
|
||||
if incidentUUID == [16]byte{} {
|
||||
return IncidentView{}, ErrIncidentNotFound
|
||||
}
|
||||
database, err := store.openDatabase()
|
||||
if err != nil {
|
||||
return IncidentView{}, err
|
||||
}
|
||||
view, err := scanIncidentView(database.QueryRowContext(ctx, `SELECT incident_uuid, detected_at, resolved_at, state, kind, scope, client_id, issue_uuid, summary, data_loss, automatically_repairable FROM storage_incidents WHERE incident_uuid = ?`, incidentUUID[:]))
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return IncidentView{}, ErrIncidentNotFound
|
||||
}
|
||||
return view, err
|
||||
}
|
||||
|
||||
func scanIncidentView(scanner interface{ Scan(...any) error }) (IncidentView, error) {
|
||||
var view IncidentView
|
||||
var encoded, issue []byte
|
||||
var detected int64
|
||||
var resolved sql.NullInt64
|
||||
var state, kind uint32
|
||||
var client sql.NullString
|
||||
var dataLoss, repairable int
|
||||
if err := scanner.Scan(&encoded, &detected, &resolved, &state, &kind, &view.Scope, &client, &issue, &view.Summary, &dataLoss, &repairable); err != nil {
|
||||
return view, err
|
||||
}
|
||||
if len(encoded) != 16 || (issue != nil && len(issue) != 16) || state < uint32(IncidentOpen) || state > uint32(IncidentAcknowledged) || kind < uint32(IncidentUncommittedTail) || kind > uint32(IncidentDiskExhaustion) {
|
||||
return view, ErrInvalidSegmentRecord
|
||||
}
|
||||
copy(view.IncidentUUID[:], encoded)
|
||||
view.DetectedAt = time.Unix(0, detected).UTC()
|
||||
view.ResolvedAt = nullableTime(resolved)
|
||||
view.State, view.Kind = IncidentState(state), IncidentKind(kind)
|
||||
if client.Valid {
|
||||
view.ClientID = client.String
|
||||
}
|
||||
view.DataLoss, view.AutomaticallyRepairable = dataLoss == 1, repairable == 1
|
||||
if issue != nil {
|
||||
value := [16]byte{}
|
||||
copy(value[:], issue)
|
||||
view.IssueUUID = &value
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
|
||||
"github.com/rvbox/rvbox/internal/domain"
|
||||
)
|
||||
|
||||
var ErrSignalDeliveryUnavailable = errors.New("signal delivery is unavailable for this command state")
|
||||
|
||||
type SignalInput struct {
|
||||
IssueUUID domain.UUID
|
||||
ClientID string
|
||||
RequestUUID domain.UUID
|
||||
Signal rvboxv1.SignalKind
|
||||
Hash [32]byte
|
||||
OccurredAt time.Time
|
||||
}
|
||||
|
||||
type SignalResult struct {
|
||||
CommandRevision uint64
|
||||
Duplicate bool
|
||||
Cancelled bool
|
||||
}
|
||||
|
||||
// SignalCommand applies the only locally-completable signal operation: a
|
||||
// queued command can be cancelled before dispatch. For dispatched/accepted/
|
||||
// running work the revisioned remote delivery path is intentionally not
|
||||
// claimed until a live session control queue is available.
|
||||
func (store *Store) SignalCommand(ctx context.Context, input SignalInput) (SignalResult, error) {
|
||||
if input.IssueUUID == (domain.UUID{}) || input.ClientID == "" || input.RequestUUID == (domain.UUID{}) || input.Signal < rvboxv1.SignalKind_SIGNAL_HUP || input.Signal > rvboxv1.SignalKind_SIGNAL_USR2 || input.Hash == [32]byte{} || input.OccurredAt.IsZero() {
|
||||
return SignalResult{}, errors.New("invalid signal request")
|
||||
}
|
||||
store.writeMu.Lock()
|
||||
defer store.writeMu.Unlock()
|
||||
database, err := store.openDatabase()
|
||||
if err != nil {
|
||||
return SignalResult{}, err
|
||||
}
|
||||
method := "signal_command"
|
||||
target := input.IssueUUID.String()
|
||||
if existing, found, lookupErr := lookupControlMutation(ctx, database, input.RequestUUID, method, target, input.Hash); lookupErr != nil {
|
||||
return SignalResult{}, lookupErr
|
||||
} else if found {
|
||||
if len(existing) != 8 {
|
||||
return SignalResult{}, ErrInvalidSegmentRecord
|
||||
}
|
||||
return SignalResult{CommandRevision: binary.BigEndian.Uint64(existing), Duplicate: true}, nil
|
||||
}
|
||||
tx, err := database.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return SignalResult{}, err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
var lifecycle uint32
|
||||
var revision uint64
|
||||
err = tx.QueryRowContext(ctx, `SELECT lifecycle, revision FROM commands WHERE issue_uuid = ? AND client_id = ?`, input.IssueUUID[:], input.ClientID).Scan(&lifecycle, &revision)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SignalResult{}, ErrCommandNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return SignalResult{}, err
|
||||
}
|
||||
if lifecycle != uint32(rvboxv1.CommandLifecycle_COMMAND_QUEUED) {
|
||||
return SignalResult{}, ErrSignalDeliveryUnavailable
|
||||
}
|
||||
if revision == ^uint64(0) {
|
||||
return SignalResult{}, errors.New("command revision exhausted")
|
||||
}
|
||||
nextRevision := revision + 1
|
||||
if _, err := tx.ExecContext(ctx, `UPDATE commands SET lifecycle = ?, terminal_time = ?, revision = ? WHERE issue_uuid = ? AND client_id = ? AND lifecycle = ? AND revision = ?`, rvboxv1.CommandLifecycle_COMMAND_CANCELLED, input.OccurredAt.UTC().UnixNano(), nextRevision, input.IssueUUID[:], input.ClientID, rvboxv1.CommandLifecycle_COMMAND_QUEUED, revision); err != nil {
|
||||
return SignalResult{}, err
|
||||
}
|
||||
result := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(result, nextRevision)
|
||||
_, err = tx.ExecContext(ctx, `INSERT INTO control_mutations (request_uuid, method, owner_kind, owner_id, target, immutable_sha256, assigned_revision, result, created_at) VALUES (?, ?, 'command', ?, ?, ?, ?, ?, ?)`, input.RequestUUID[:], method, input.IssueUUID.String(), target, input.Hash[:], nextRevision, result, input.OccurredAt.UTC().UnixNano())
|
||||
if err != nil {
|
||||
return SignalResult{}, err
|
||||
}
|
||||
// Keep a compact audit digest; command text/stdin is never copied to the
|
||||
// audit payload.
|
||||
auditPayload := []byte{byte(input.Signal)}
|
||||
auditHash := sha256.Sum256(auditPayload)
|
||||
_, err = tx.ExecContext(ctx, `INSERT INTO audit_events (occurred_at, source, action, outcome, compression, payload, raw_bytes, stored_bytes, sha256) VALUES (?, 'control', ?, 'success', 1, ?, ?, ?, ?)`, input.OccurredAt.UTC().UnixNano(), method, auditPayload, len(auditPayload), len(auditPayload), auditHash[:])
|
||||
if err != nil {
|
||||
return SignalResult{}, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return SignalResult{}, err
|
||||
}
|
||||
return SignalResult{CommandRevision: nextRevision, Cancelled: true}, nil
|
||||
}
|
||||
@@ -0,0 +1,163 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"math"
|
||||
"time"
|
||||
|
||||
"github.com/rvbox/rvbox/internal/domain"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrCommandTerminal = errors.New("command is already terminal")
|
||||
ErrStdinConflict = errors.New("stdin request conflicts with an earlier request")
|
||||
)
|
||||
|
||||
type StdinWriteInput struct {
|
||||
IssueUUID domain.UUID
|
||||
ClientID string
|
||||
RequestUUID domain.UUID
|
||||
Data []byte
|
||||
AppendNewline bool
|
||||
Close bool
|
||||
ImmutableHash [32]byte
|
||||
OccurredAt time.Time
|
||||
}
|
||||
|
||||
type StdinWriteResult struct {
|
||||
WriteSeq uint64
|
||||
Duplicate bool
|
||||
}
|
||||
|
||||
// AppendStdin durably records one ordered stdin intent. It does not claim
|
||||
// delivery to a process; an agent acknowledgement is a later protocol event.
|
||||
// Reusing RequestUUID with the same immutable hash returns the original write
|
||||
// sequence without adding a second row.
|
||||
func (store *Store) AppendStdin(ctx context.Context, input StdinWriteInput) (StdinWriteResult, error) {
|
||||
return store.appendStdin(ctx, input)
|
||||
}
|
||||
|
||||
func (store *Store) CloseStdin(ctx context.Context, input StdinWriteInput) (StdinWriteResult, error) {
|
||||
input.Close = true
|
||||
return store.appendStdin(ctx, input)
|
||||
}
|
||||
|
||||
func (store *Store) appendStdin(ctx context.Context, input StdinWriteInput) (StdinWriteResult, error) {
|
||||
if input.IssueUUID == (domain.UUID{}) || input.ClientID == "" || input.RequestUUID == (domain.UUID{}) || input.OccurredAt.IsZero() || input.ImmutableHash == [32]byte{} || (!input.Close && len(input.Data) == 0) || (input.Close && len(input.Data) != 0) || len(input.Data) > 64<<10 {
|
||||
return StdinWriteResult{}, errors.New("invalid stdin write")
|
||||
}
|
||||
stored, err := compressCommandSpec(input.Data)
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
charge, err := EstimateCharge(ChargeInput{EncodedBytes: uint64(len(stored)), SQLiteRows: 1, IndexEntries: 1})
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
method := "append_stdin"
|
||||
if input.Close {
|
||||
method = "close_stdin"
|
||||
}
|
||||
target := input.IssueUUID.String()
|
||||
store.writeMu.Lock()
|
||||
defer store.writeMu.Unlock()
|
||||
database, err := store.openDatabase()
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
if existing, found, lookupErr := lookupControlMutation(ctx, database, input.RequestUUID, method, target, input.ImmutableHash); lookupErr != nil {
|
||||
return StdinWriteResult{}, lookupErr
|
||||
} else if found {
|
||||
if len(existing) != 8 {
|
||||
return StdinWriteResult{}, ErrInvalidSegmentRecord
|
||||
}
|
||||
return StdinWriteResult{WriteSeq: binary.BigEndian.Uint64(existing), Duplicate: true}, nil
|
||||
}
|
||||
var lifecycle uint32
|
||||
var commandCharged, closeout, clientCharged, serverCharged uint64
|
||||
err = database.QueryRowContext(ctx, `SELECT commands.lifecycle, commands.charged_bytes, commands.closeout_remaining_bytes,
|
||||
clients.charged_bytes, storage_counters.command_charged_bytes
|
||||
FROM commands JOIN clients ON clients.client_id = commands.client_id
|
||||
JOIN storage_counters ON storage_counters.singleton = 1
|
||||
WHERE commands.issue_uuid = ? AND commands.client_id = ?`, input.IssueUUID[:], input.ClientID).Scan(&lifecycle, &commandCharged, &closeout, &clientCharged, &serverCharged)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return StdinWriteResult{}, ErrCommandNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
if lifecycle < 1 || lifecycle > 4 {
|
||||
return StdinWriteResult{}, ErrCommandTerminal
|
||||
}
|
||||
freeBytes, err := store.freeSpaceProbe.AvailableBytes(store.dataDir)
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
reservation, err := CheckReservation(store.quotaLimits, ReservationState{CommandTotalCharged: commandCharged, CloseoutRemaining: closeout, ClientTotalCharged: clientCharged, ServerTotalCharged: serverCharged, FilesystemFreeBytes: freeBytes}, ReservationRequest{ChargedBytes: charge, PhysicalBytes: uint64(len(stored))})
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
var writeSeq uint64
|
||||
if err := database.QueryRowContext(ctx, `SELECT COALESCE(MAX(write_seq), 0) + 1 FROM stdin_writes WHERE issue_uuid = ?`, input.IssueUUID[:]).Scan(&writeSeq); err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
if writeSeq == 0 || writeSeq > math.MaxInt64 {
|
||||
return StdinWriteResult{}, errors.New("stdin write sequence exhausted")
|
||||
}
|
||||
digest := sha256.Sum256(input.Data)
|
||||
appendNewline := 0
|
||||
if input.AppendNewline {
|
||||
appendNewline = 1
|
||||
}
|
||||
closeIntent := 0
|
||||
if input.Close {
|
||||
closeIntent = 1
|
||||
}
|
||||
tx, err := database.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
_, err = tx.ExecContext(ctx, `INSERT INTO stdin_writes (issue_uuid, write_seq, payload, raw_bytes, stored_bytes, compression, sha256, append_newline, close_intent, acknowledged) VALUES (?, ?, ?, ?, ?, 2, ?, ?, ?, 0)`, input.IssueUUID[:], writeSeq, stored, len(input.Data), len(stored), digest[:], appendNewline, closeIntent)
|
||||
if err == nil {
|
||||
_, err = tx.ExecContext(ctx, `UPDATE commands SET charged_bytes = ?, closeout_remaining_bytes = ? WHERE issue_uuid = ? AND charged_bytes = ?`, reservation.CommandTotalCharged, reservation.CloseoutRemaining, input.IssueUUID[:], commandCharged)
|
||||
}
|
||||
if err == nil {
|
||||
_, err = tx.ExecContext(ctx, `UPDATE clients SET charged_bytes = ? WHERE client_id = ? AND charged_bytes = ?`, reservation.ClientTotalCharged, input.ClientID, clientCharged)
|
||||
}
|
||||
if err == nil {
|
||||
_, err = tx.ExecContext(ctx, `UPDATE storage_counters SET command_charged_bytes = ? WHERE singleton = 1 AND command_charged_bytes = ?`, reservation.ServerTotalCharged, serverCharged)
|
||||
}
|
||||
result := make([]byte, 8)
|
||||
if err == nil {
|
||||
binary.BigEndian.PutUint64(result, writeSeq)
|
||||
_, err = tx.ExecContext(ctx, `INSERT INTO control_mutations (request_uuid, method, owner_kind, owner_id, target, immutable_sha256, assigned_write_seq, result, created_at) VALUES (?, ?, 'command', ?, ?, ?, ?, ?, ?)`, input.RequestUUID[:], method, input.IssueUUID.String(), target, input.ImmutableHash[:], writeSeq, result, input.OccurredAt.UTC().UnixNano())
|
||||
}
|
||||
if err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return StdinWriteResult{}, err
|
||||
}
|
||||
return StdinWriteResult{WriteSeq: writeSeq}, nil
|
||||
}
|
||||
|
||||
func lookupControlMutation(ctx context.Context, database *sql.DB, requestUUID domain.UUID, method, target string, hash [32]byte) ([]byte, bool, error) {
|
||||
var storedMethod, storedTarget string
|
||||
var storedHash, result []byte
|
||||
err := database.QueryRowContext(ctx, `SELECT method, target, immutable_sha256, result FROM control_mutations WHERE request_uuid = ?`, requestUUID[:]).Scan(&storedMethod, &storedTarget, &storedHash, &result)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
if storedMethod != method || storedTarget != target || len(storedHash) != 32 || string(storedHash) != string(hash[:]) {
|
||||
return nil, false, ErrMutationConflict
|
||||
}
|
||||
return append([]byte(nil), result...), true, nil
|
||||
}
|
||||
@@ -59,6 +59,36 @@ layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["internal/server/control/listen_unix_test.go:TestListenUnixRefusesNonSocketPath_BH_CONTROL_05"]
|
||||
|
||||
[[requirements]]
|
||||
id = "HP-CTL-10"
|
||||
layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["internal/server/control/service_test.go:TestFollowCommandStreamsDurableEvents_HP_CONTROL_10"]
|
||||
|
||||
[[requirements]]
|
||||
id = "HP-CTL-11"
|
||||
layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["cmd/rvc/main_test.go:TestGlobalSocketAndCLIValueParsing_HP_CTL_11"]
|
||||
|
||||
[[requirements]]
|
||||
id = "HP-CTL-12"
|
||||
layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["internal/server/control/service_test.go:TestGetOutputSlicesWithAuthenticatedCursor_HP_CONTROL_12"]
|
||||
|
||||
[[requirements]]
|
||||
id = "BH-CTL-13"
|
||||
layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["internal/server/control/service_test.go:TestControlMutationIdempotencyAndQueuedCancellation_BH_CONTROL_13"]
|
||||
|
||||
[[requirements]]
|
||||
id = "HP-CTL-14"
|
||||
layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["internal/server/control/service_test.go:TestStorageIncidentControlLifecycle_HP_CONTROL_14"]
|
||||
|
||||
[[requirements]]
|
||||
id = "BH-CTL-01"
|
||||
layer = "unit"
|
||||
@@ -278,6 +308,12 @@ layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["internal/client/agent/handshake_test.go:TestApplyEventAckValidatesAndReleasesPrefix_HP_EVENT_03"]
|
||||
|
||||
[[requirements]]
|
||||
id = "HP-EVENT-04"
|
||||
layer = "unit"
|
||||
status = "implemented"
|
||||
tests = ["internal/server/store/events_read_test.go:TestReadCommandEventsValidatesSegmentAndPreservesOrder_HP_EVENT_04"]
|
||||
|
||||
[[requirements]]
|
||||
id = "HP-SES-05"
|
||||
layer = "integration"
|
||||
|
||||
Reference in New Issue
Block a user