From 8df56263fb4aec1aed54a88d8ccfe659b81a50e0 Mon Sep 17 00:00:00 2001 From: cabbage Date: Sun, 6 Sep 2026 10:18:37 +0000 Subject: [PATCH] feat: implement control history and mutation APIs --- cmd/rvbox-server/main.go | 2 +- cmd/rvc/main.go | 252 ++++++++++++- cmd/rvc/main_test.go | 34 ++ internal/server/control/service.go | 437 +++++++++++++++++++++- internal/server/control/service_test.go | 170 +++++++++ internal/server/store/events_read.go | 156 ++++++++ internal/server/store/events_read_test.go | 52 +++ internal/server/store/incidents_read.go | 118 ++++++ internal/server/store/signal.go | 98 +++++ internal/server/store/stdin.go | 163 ++++++++ test/coverage.toml | 36 ++ 11 files changed, 1515 insertions(+), 3 deletions(-) create mode 100644 cmd/rvc/main_test.go create mode 100644 internal/server/store/events_read.go create mode 100644 internal/server/store/events_read_test.go create mode 100644 internal/server/store/incidents_read.go create mode 100644 internal/server/store/signal.go create mode 100644 internal/server/store/stdin.go diff --git a/cmd/rvbox-server/main.go b/cmd/rvbox-server/main.go index c5184bc..45350e9 100644 --- a/cmd/rvbox-server/main.go +++ b/cmd/rvbox-server/main.go @@ -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}, }) diff --git a/cmd/rvc/main.go b/cmd/rvc/main.go index bc78a7a..d78ce24 100644 --- a/cmd/rvc/main.go +++ b/cmd/rvc/main.go @@ -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,15 +235,254 @@ 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 { return nil diff --git a/cmd/rvc/main_test.go b/cmd/rvc/main_test.go new file mode 100644 index 0000000..cda7a03 --- /dev/null +++ b/cmd/rvc/main_test.go @@ -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) + } +} diff --git a/internal/server/control/service.go b/internal/server/control/service.go index 8b90b77..14e1093 100644 --- a/internal/server/control/service.go +++ b/internal/server/control/service.go @@ -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: diff --git a/internal/server/control/service_test.go b/internal/server/control/service_test.go index 479d496..ef2ba9b 100644 --- a/internal/server/control/service_test.go +++ b/internal/server/control/service_test.go @@ -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}) diff --git a/internal/server/store/events_read.go b/internal/server/store/events_read.go new file mode 100644 index 0000000..2405378 --- /dev/null +++ b/internal/server/store/events_read.go @@ -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 +} diff --git a/internal/server/store/events_read_test.go b/internal/server/store/events_read_test.go new file mode 100644 index 0000000..b5af8cd --- /dev/null +++ b/internal/server/store/events_read_test.go @@ -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") + } +} diff --git a/internal/server/store/incidents_read.go b/internal/server/store/incidents_read.go new file mode 100644 index 0000000..bf3d28d --- /dev/null +++ b/internal/server/store/incidents_read.go @@ -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 +} diff --git a/internal/server/store/signal.go b/internal/server/store/signal.go new file mode 100644 index 0000000..2009a5d --- /dev/null +++ b/internal/server/store/signal.go @@ -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 +} diff --git a/internal/server/store/stdin.go b/internal/server/store/stdin.go new file mode 100644 index 0000000..a773f9e --- /dev/null +++ b/internal/server/store/stdin.go @@ -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 +} diff --git a/test/coverage.toml b/test/coverage.toml index 2bc9176..a41c452 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -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"