diff --git a/cmd/rvbox-server/main.go b/cmd/rvbox-server/main.go index 36e03cb..3a52604 100644 --- a/cmd/rvbox-server/main.go +++ b/cmd/rvbox-server/main.go @@ -13,10 +13,13 @@ import ( "os/signal" "syscall" + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/agentproto" "github.com/rvbox/rvbox/internal/config" + "github.com/rvbox/rvbox/internal/server/control" "github.com/rvbox/rvbox/internal/server/session" "github.com/rvbox/rvbox/internal/server/store" + "google.golang.org/grpc" ) func main() { @@ -60,6 +63,21 @@ func run(configPath string) error { if err != nil { return fmt.Errorf("listen for agents: %w", err) } + defer listener.Close() + controlListener, cleanupControl, err := control.ListenUnix(configured.Server.ControlSocket) + if err != nil { + return err + } + defer func() { _ = cleanupControl() }() + controlService, err := control.NewService(control.Options{ + Store: persistence, DefaultQueueTTL: configured.Queue.DefaultTTL, DefaultQueueTTLSet: true, + Limits: agentproto.Limits{MaxEnvelopeBytes: configured.Protocol.MaxAgentEnvelopeBytes, MaxExecutionSpecBytes: configured.Protocol.MaxExecutionSpecBytes, MaxRawChunkBytes: configured.Protocol.MaxRawChunkBytes, MaxScriptBytes: configured.Protocol.MaxScriptBytes, MaxDetailBytes: configured.Storage.ProtocolDetailMaxBytes}, + }) + if err != nil { + return err + } + grpcServer := grpc.NewServer(grpc.MaxRecvMsgSize(int(configured.Protocol.MaxControlRequestBytes)), grpc.MaxSendMsgSize(int(configured.Protocol.MaxControlRequestBytes))) + rvboxv1.RegisterControlServer(grpcServer, controlService) agent := &session.AgentServer{ Store: persistence, Registry: session.NewRegistry(), Path: configured.Server.AgentPath, Limits: agentproto.Limits{ @@ -73,8 +91,9 @@ func run(configPath string) error { HeartbeatIdle: configured.Protocol.HeartbeatIdle, LivenessTimeout: configured.Protocol.LivenessTimeout, } httpServer := &http.Server{Handler: agent, ReadHeaderTimeout: configured.Flow.WriteDeadline} - serveError := make(chan error, 1) + serveError := make(chan error, 2) go func() { serveError <- httpServer.Serve(listener) }() + go func() { serveError <- grpcServer.Serve(controlListener) }() signals := make(chan os.Signal, 1) signal.Notify(signals, os.Interrupt, syscall.SIGTERM) @@ -82,12 +101,22 @@ func run(configPath string) error { select { case err := <-serveError: if errors.Is(err, http.ErrServerClosed) { + grpcServer.Stop() return nil } + grpcServer.Stop() return err case <-signals: shutdownContext, cancel := context.WithTimeout(context.Background(), configured.Server.ShutdownGrace) defer cancel() - return httpServer.Shutdown(shutdownContext) + httpErr := httpServer.Shutdown(shutdownContext) + grpcDone := make(chan struct{}) + go func() { grpcServer.GracefulStop(); close(grpcDone) }() + select { + case <-grpcDone: + case <-shutdownContext.Done(): + grpcServer.Stop() + } + return httpErr } } diff --git a/cmd/rvc/main.go b/cmd/rvc/main.go index c779a96..bc78a7a 100644 --- a/cmd/rvc/main.go +++ b/cmd/rvc/main.go @@ -1,4 +1,287 @@ -// Command rvc is the RVBox control-plane command-line client. +// Command rvc is the RVBox local control-plane CLI. It never opens the server +// database; all state changes and reads go through the Unix gRPC socket. package main -func main() {} +import ( + "context" + "crypto/sha256" + "errors" + "flag" + "fmt" + "io" + "net" + "os" + "path/filepath" + "strings" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/protobuf/types/known/durationpb" +) + +const defaultControlSocket = "/run/rvbox/server.sock" + +func main() { + if err := run(os.Args[1:], os.Stdout, os.Stderr); err != nil { + fmt.Fprintln(os.Stderr, "rvc:", err) + os.Exit(1) + } +} + +func run(args []string, output, diagnostics io.Writer) error { + if len(args) == 0 { + return errors.New("a command is required (stat or run)") + } + socket, args, err := globalSocket(args) + if err != nil { + return err + } + if len(args) == 0 { + return errors.New("a command is required (stat or run)") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + connection, err := dial(ctx, socket) + if err != nil { + return err + } + defer connection.Close() + client := rvboxv1.NewControlClient(connection) + switch args[0] { + case "stat": + return stat(ctx, client, args[1:], output) + case "run": + return runCommand(ctx, client, args[1:], output, diagnostics) + default: + return fmt.Errorf("unknown command %q", args[0]) + } +} + +func globalSocket(args []string) (string, []string, error) { + socket := defaultControlSocket + remaining := make([]string, 0, len(args)) + for index := 0; index < len(args); index++ { + if args[index] == "--socket" { + if index+1 >= len(args) { + return "", nil, errors.New("--socket requires a path") + } + socket = args[index+1] + index++ + continue + } + if strings.HasPrefix(args[index], "--socket=") { + socket = strings.TrimPrefix(args[index], "--socket=") + continue + } + remaining = append(remaining, args[index]) + } + if socket == "" || !filepath.IsAbs(socket) { + return "", nil, errors.New("--socket must be an absolute path") + } + return socket, remaining, nil +} + +func dial(ctx context.Context, socket string) (*grpc.ClientConn, error) { + dialer := func(ctx context.Context, _ string) (net.Conn, error) { + return (&net.Dialer{}).DialContext(ctx, "unix", socket) + } + return grpc.DialContext(ctx, "passthrough:///rvbox-control", grpc.WithTransportCredentials(insecure.NewCredentials()), grpc.WithContextDialer(dialer), grpc.WithBlock()) +} + +func stat(ctx context.Context, client rvboxv1.ControlClient, args []string, output io.Writer) error { + flags := flag.NewFlagSet("stat", flag.ContinueOnError) + flags.SetOutput(io.Discard) + all := flags.Bool("all", false, "show all command pages") + pageSize := flags.Uint("per-page", 100, "number of records per page") + if err := flags.Parse(args); err != nil { + return err + } + positionals := flags.Args() + if len(positionals) == 0 { + var token string + for { + response, err := client.ListClients(ctx, &rvboxv1.ListClientsRequest{PageSize: uint32(*pageSize), PageToken: token}) + if err != nil { + return err + } + for _, item := range response.GetClients() { + fmt.Fprintf(output, "client %s connected=%t platform=%s running=%d queued=%d\n", item.GetClientId(), item.GetConnected(), item.GetPlatform(), item.GetRunningCommands(), item.GetQueuedCommands()) + } + if !*all || response.GetNextPageToken() == "" { + return nil + } + token = response.GetNextPageToken() + } + } + if len(positionals) > 2 { + return errors.New("stat accepts at most CLIENT and ISSUE_UUID") + } + clientID := positionals[0] + if len(positionals) == 1 { + response, err := client.GetClient(ctx, &rvboxv1.GetClientRequest{ClientId: clientID}) + if err != nil { + return err + } + item := response.GetClient() + fmt.Fprintf(output, "client %s connected=%t platform=%s running=%d queued=%d instance=%s\n", item.GetClientId(), item.GetConnected(), item.GetPlatform(), item.GetRunningCommands(), item.GetQueuedCommands(), item.GetClientInstanceId()) + if item.GetPendingInstanceId() != "" { + fmt.Fprintf(output, "pending_instance=%s\n", item.GetPendingInstanceId()) + } + return nil + } + response, err := client.GetCommand(ctx, &rvboxv1.GetCommandRequest{ClientId: clientID, IssueUuid: positionals[1]}) + if err != nil { + return err + } + item := response.GetCommand() + fmt.Fprintf(output, "command %s client=%s lifecycle=%s revision=%d events=%d\n", item.GetIssueUuid(), item.GetTargetClientId(), item.GetLifecycle(), item.GetCommandRevision(), item.GetLastEventSeq()) + if item.GetQueueExpiryTime() != nil { + fmt.Fprintf(output, "queue_expiry=%s\n", item.GetQueueExpiryTime().AsTime().UTC().Format(time.RFC3339Nano)) + } + if item.GetTerminalTime() != nil { + fmt.Fprintf(output, "terminal_time=%s\n", item.GetTerminalTime().AsTime().UTC().Format(time.RFC3339Nano)) + } + return nil +} + +func runCommand(ctx context.Context, client rvboxv1.ControlClient, args []string, output, diagnostics io.Writer) error { + flags := flag.NewFlagSet("run", flag.ContinueOnError) + flags.SetOutput(io.Discard) + background := flags.Bool("background", false, "return after durable admission") + cwd := flags.String("cwd", "", "command working directory") + shell := flags.String("shell", "", "shell (sh, bash, cmd, powershell)") + requestID := flags.String("request-id", "", "canonical UUIDv7 used for idempotent admission") + queueTTL := flags.Duration("queue-ttl", -1, "queue TTL; zero means no expiry") + scriptPath := flags.String("script", "", "script file path") + envValues := repeatedFlag{} + profileValues := repeatedFlag{} + flags.Var(&envValues, "env", "environment override KEY=VALUE") + flags.Var(&profileValues, "profile", "execution profile") + if err := flags.Parse(args); err != nil { + return err + } + positionals := flags.Args() + if len(positionals) == 0 { + return errors.New("run requires CLIENT and COMMAND (or --script PATH CLIENT)") + } + if *requestID == "" { + generated, err := domain.NewUUIDv7() + if err != nil { + return err + } + *requestID = generated.String() + } + if _, err := domain.ParseUUIDv7(*requestID); err != nil { + return fmt.Errorf("--request-id: %w", err) + } + spec := &rvboxv1.ExecutionSpec{Cwd: *cwd} + spec.ShellType, _ = parseShell(*shell) + for _, value := range envValues { + parts := strings.SplitN(value, "=", 2) + if len(parts) != 2 || parts[0] == "" { + return fmt.Errorf("--env must be KEY=VALUE: %q", value) + } + if spec.EnvOverrides == nil { + spec.EnvOverrides = make(map[string]string) + } + spec.EnvOverrides[parts[0]] = parts[1] + } + for _, value := range profileValues { + profile, err := parseProfile(value) + if err != nil { + return err + } + spec.ExecutionProfiles = append(spec.ExecutionProfiles, profile) + } + var clientID string + if *scriptPath != "" { + if len(positionals) != 1 { + return errors.New("script form accepts exactly CLIENT") + } + clientID = positionals[0] + data, err := os.ReadFile(*scriptPath) + if err != nil { + return err + } + if len(data) > 10<<20 { + return errors.New("script exceeds the 10 MiB limit") + } + digest := sha256.Sum256(data) + spec.Source = &rvboxv1.ExecutionSpec_Script{Script: &rvboxv1.ScriptDescriptor{Filename: filepath.Base(*scriptPath), SizeBytes: uint64(len(data)), Sha256: digest[:]}} + request := &rvboxv1.RunCommandRequest{TargetClientId: clientID, Spec: spec, ScriptContent: data, RequestId: *requestID, QueueTtl: queueDuration(*queueTTL)} + return printRunResponse(ctx, client, request, output, *background) + } + if len(positionals) < 2 { + return errors.New("run requires CLIENT and COMMAND") + } + clientID = positionals[0] + spec.Source = &rvboxv1.ExecutionSpec_CommandText{CommandText: strings.Join(positionals[1:], " ")} + request := &rvboxv1.RunCommandRequest{TargetClientId: clientID, Spec: spec, RequestId: *requestID, QueueTtl: queueDuration(*queueTTL)} + _ = diagnostics + return printRunResponse(ctx, client, request, output, *background) +} + +func printRunResponse(ctx context.Context, client rvboxv1.ControlClient, request *rvboxv1.RunCommandRequest, output io.Writer, _ bool) error { + response, err := client.RunCommand(ctx, request) + if err != nil { + return err + } + fmt.Fprintf(output, "%s %s\n", response.GetIssueUuid(), response.GetLifecycle()) + return nil +} + +func queueDuration(value time.Duration) *durationpb.Duration { + if value < 0 { + return nil + } + return durationpb.New(value) +} + +func parseShell(value string) (rvboxv1.ShellType, error) { + switch strings.ToLower(value) { + case "", "default": + return rvboxv1.ShellType_SHELL_TYPE_UNSPECIFIED, nil + case "sh": + return rvboxv1.ShellType_SHELL_SH, nil + case "bash": + return rvboxv1.ShellType_SHELL_BASH, nil + case "cmd": + return rvboxv1.ShellType_SHELL_CMD, nil + case "powershell", "pwsh": + return rvboxv1.ShellType_SHELL_POWERSHELL, nil + default: + return rvboxv1.ShellType_SHELL_TYPE_UNSPECIFIED, fmt.Errorf("unknown shell %q", value) + } +} + +func parseProfile(value string) (rvboxv1.ExecutionProfile, error) { + switch strings.ToLower(value) { + case "light": + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_LIGHT, nil + case "cpu-medium": + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_MEDIUM, nil + case "cpu-heavy": + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_CPU_HEAVY, nil + case "mem-medium": + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_MEM_MEDIUM, nil + case "mem-heavy": + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_MEM_HEAVY, nil + case "disk-medium": + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_DISK_MEDIUM, nil + case "disk-heavy": + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_DISK_HEAVY, nil + default: + return rvboxv1.ExecutionProfile_EXECUTION_PROFILE_UNSPECIFIED, fmt.Errorf("unknown profile %q", value) + } +} + +type repeatedFlag []string + +func (flag *repeatedFlag) String() string { return strings.Join(*flag, ",") } +func (flag *repeatedFlag) Set(value string) error { + *flag = append(*flag, value) + return nil +} diff --git a/internal/domain/cursor.go b/internal/domain/cursor.go index 26b3b45..b7276e5 100644 --- a/internal/domain/cursor.go +++ b/internal/domain/cursor.go @@ -26,6 +26,7 @@ const ( CursorKindCommands CursorKind = iota + 1 CursorKindOutput CursorKindIncidents + CursorKindClients ) type Cursor struct { @@ -108,5 +109,5 @@ func (codec *CursorCodec) Decode(token string, expectedKind CursorKind, expected } func validCursorKind(kind CursorKind) bool { - return kind >= CursorKindCommands && kind <= CursorKindIncidents + return kind >= CursorKindCommands && kind <= CursorKindClients } diff --git a/internal/server/control/listen_unix.go b/internal/server/control/listen_unix.go new file mode 100644 index 0000000..ecddda2 --- /dev/null +++ b/internal/server/control/listen_unix.go @@ -0,0 +1,80 @@ +//go:build !windows + +package control + +import ( + "errors" + "fmt" + "net" + "os" + "path/filepath" + "syscall" + "time" +) + +var ErrControlSocketInUse = errors.New("control socket is already in use") + +// ListenUnix creates a private local control socket and returns a cleanup +// function that removes only the socket created by this process. An existing +// path is removed only after it is proven to be a socket owned by this UID and +// a connection probe fails. +func ListenUnix(path string) (net.Listener, func() error, error) { + if path == "" || !filepath.IsAbs(path) || filepath.Clean(path) == string(filepath.Separator) { + return nil, nil, errors.New("control socket path must be an absolute non-root path") + } + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + return nil, nil, fmt.Errorf("create control socket directory: %w", err) + } + if info, err := os.Lstat(path); err == nil { + if info.Mode()&os.ModeSymlink != 0 || info.Mode()&os.ModeSocket == 0 || info.Mode().Perm()&0o077 != 0 || !ownedByCurrentUser(info) { + return nil, nil, errors.New("existing control socket path is not a private owned socket") + } + probe, probeErr := net.DialTimeout("unix", path, 100*time.Millisecond) + if probeErr == nil { + _ = probe.Close() + return nil, nil, ErrControlSocketInUse + } + if err := os.Remove(path); err != nil { + return nil, nil, fmt.Errorf("remove stale control socket: %w", err) + } + } else if !errors.Is(err, os.ErrNotExist) { + return nil, nil, fmt.Errorf("inspect control socket: %w", err) + } + listener, err := net.Listen("unix", path) + if err != nil { + return nil, nil, fmt.Errorf("listen on control socket: %w", err) + } + if err := os.Chmod(path, 0o600); err != nil { + _ = listener.Close() + _ = os.Remove(path) + return nil, nil, fmt.Errorf("protect control socket: %w", err) + } + createdInfo, err := os.Lstat(path) + if err != nil { + _ = listener.Close() + _ = os.Remove(path) + return nil, nil, fmt.Errorf("stat control socket: %w", err) + } + cleanup := func() error { + if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + return err + } + info, err := os.Lstat(path) + if errors.Is(err, os.ErrNotExist) { + return nil + } + if err != nil { + return err + } + if info.Mode()&os.ModeSocket == 0 || info.Mode().Perm()&0o077 != 0 || !ownedByCurrentUser(info) || !os.SameFile(createdInfo, info) { + return errors.New("refusing to remove changed control socket") + } + return os.Remove(path) + } + return listener, cleanup, nil +} + +func ownedByCurrentUser(info os.FileInfo) bool { + stat, ok := info.Sys().(*syscall.Stat_t) + return ok && uint32(stat.Uid) == uint32(os.Getuid()) +} diff --git a/internal/server/control/listen_unix_test.go b/internal/server/control/listen_unix_test.go new file mode 100644 index 0000000..e6b6977 --- /dev/null +++ b/internal/server/control/listen_unix_test.go @@ -0,0 +1,52 @@ +//go:build !windows + +package control + +import ( + "net" + "os" + "path/filepath" + "testing" +) + +func TestListenUnixProtectsAndCleansOwnedSocket_HP_CONTROL_04(t *testing.T) { + path := filepath.Join(t.TempDir(), "control.sock") + listener, cleanup, err := ListenUnix(path) + if err != nil { + t.Fatal(err) + } + defer listener.Close() + info, err := os.Stat(path) + if err != nil { + t.Fatal(err) + } + if info.Mode()&os.ModeSocket == 0 || info.Mode().Perm() != 0o600 { + t.Fatalf("socket mode/type = %v/%o", info.Mode(), info.Mode().Perm()) + } + if _, err := net.Dial("unix", path); err != nil { + t.Fatalf("dial active socket: %v", err) + } + if _, _, err := ListenUnix(path); err != ErrControlSocketInUse { + t.Fatalf("active socket error = %v", err) + } + if err := cleanup(); err != nil { + t.Fatal(err) + } + if _, err := os.Lstat(path); !os.IsNotExist(err) { + t.Fatalf("socket after cleanup = %v", err) + } +} + +func TestListenUnixRefusesNonSocketPath_BH_CONTROL_05(t *testing.T) { + path := filepath.Join(t.TempDir(), "control.sock") + if err := os.WriteFile(path, []byte("sentinel"), 0o600); err != nil { + t.Fatal(err) + } + if _, _, err := ListenUnix(path); err == nil { + t.Fatal("regular file path accepted as control socket") + } + data, err := os.ReadFile(path) + if err != nil || string(data) != "sentinel" { + t.Fatalf("regular file changed: %q, %v", data, err) + } +} diff --git a/internal/server/control/listen_windows.go b/internal/server/control/listen_windows.go new file mode 100644 index 0000000..ad72fcf --- /dev/null +++ b/internal/server/control/listen_windows.go @@ -0,0 +1,14 @@ +//go:build windows + +package control + +import ( + "errors" + "net" +) + +var ErrControlSocketInUse = errors.New("Unix control sockets are unavailable on Windows") + +func ListenUnix(string) (net.Listener, func() error, error) { + return nil, nil, ErrControlSocketInUse +} diff --git a/internal/server/control/service.go b/internal/server/control/service.go new file mode 100644 index 0000000..288da0a --- /dev/null +++ b/internal/server/control/service.go @@ -0,0 +1,406 @@ +// Package control implements the local gRPC control plane. It is deliberately +// independent of the WebSocket session implementation: all durable decisions +// go through store APIs, while a later dispatcher may subscribe to queued work. +package control + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "encoding/binary" + "errors" + "fmt" + "math" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "github.com/rvbox/rvbox/internal/domain" + "github.com/rvbox/rvbox/internal/server/store" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/durationpb" + "google.golang.org/protobuf/types/known/timestamppb" +) + +const ( + defaultPageSize = 100 + maxPageSize = 1000 + defaultQueueTTL = 15 * time.Minute + clientsCursorVersion = "clients-v1" +) + +// Options controls the non-transport dependencies of Service. CursorKey is +// process-local by default, which invalidates old cursors after a restart +// rather than accepting a token whose query snapshot is no longer meaningful. +type Options struct { + Store *store.Store + DefaultQueueTTL time.Duration + // DefaultQueueTTLSet distinguishes an explicitly configured zero (the + // documented indefinite-queue setting) from an omitted option in tests or + // embedders that want the compiled 15-minute default. + DefaultQueueTTLSet bool + Limits agentproto.Limits + Now func() time.Time + CursorKey []byte +} + +type Service struct { + rvboxv1.UnimplementedControlServer + store *store.Store + defaultQueueTTL time.Duration + limits agentproto.Limits + now func() time.Time + cursors *domain.CursorCodec +} + +func NewService(options Options) (*Service, error) { + if options.Store == nil { + return nil, errors.New("control service requires a store") + } + if options.DefaultQueueTTL < 0 { + return nil, errors.New("default queue TTL must be non-negative") + } + if !options.DefaultQueueTTLSet && options.DefaultQueueTTL == 0 { + options.DefaultQueueTTL = defaultQueueTTL + } + if options.Limits.MaxEnvelopeBytes == 0 { + options.Limits = agentproto.DefaultLimits() + } + if options.Now == nil { + options.Now = time.Now + } + key := append([]byte(nil), options.CursorKey...) + if len(key) == 0 { + key = make([]byte, sha256.Size) + if _, err := rand.Read(key); err != nil { + return nil, fmt.Errorf("generate cursor key: %w", err) + } + } + codec, err := domain.NewCursorCodec(key) + if err != nil { + return nil, err + } + return &Service{store: options.Store, defaultQueueTTL: options.DefaultQueueTTL, limits: options.Limits, now: options.Now, cursors: codec}, nil +} + +func (service *Service) ListClients(ctx context.Context, request *rvboxv1.ListClientsRequest) (*rvboxv1.ListClientsResponse, error) { + if request == nil { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "request is required") + } + pageSize, err := pageSize(request.GetPageSize()) + if err != nil { + return nil, err + } + filter := domain.HashCursorFilters([]byte(clientsCursorVersion)) + after := "" + if request.GetPageToken() != "" { + cursor, decodeErr := service.cursors.Decode(request.GetPageToken(), domain.CursorKindClients, filter) + if decodeErr != nil || string(cursor.SnapshotBoundary) != clientsCursorVersion { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "invalid client page token") + } + after = string(cursor.Position) + } + clients, hasNext, err := service.store.ListClientViews(ctx, after, pageSize) + if err != nil { + return nil, mapStoreError(err) + } + response := &rvboxv1.ListClientsResponse{Clients: make([]*rvboxv1.ClientSummary, 0, len(clients))} + for _, client := range clients { + converted, convertErr := clientSummary(client) + if convertErr != nil { + return nil, mapStoreError(convertErr) + } + response.Clients = append(response.Clients, converted) + } + if hasNext && len(clients) > 0 { + response.NextPageToken, err = service.cursors.Encode(domain.Cursor{Kind: domain.CursorKindClients, FilterHash: filter, Position: []byte(clients[len(clients)-1].ClientID), SnapshotBoundary: []byte(clientsCursorVersion)}) + if err != nil { + return nil, mapStoreError(err) + } + } + return response, nil +} + +func (service *Service) GetClient(ctx context.Context, request *rvboxv1.GetClientRequest) (*rvboxv1.GetClientResponse, error) { + if request == nil || request.GetClientId() == "" { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id is required") + } + client, err := service.store.GetClientView(ctx, request.GetClientId()) + if err != nil { + return nil, mapStoreError(err) + } + summary, err := clientSummary(client) + if err != nil { + return nil, mapStoreError(err) + } + return &rvboxv1.GetClientResponse{Client: summary}, nil +} + +func (service *Service) ListCommands(ctx context.Context, request *rvboxv1.ListCommandsRequest) (*rvboxv1.ListCommandsResponse, error) { + if request == nil || request.GetClientId() == "" { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "client_id is required") + } + if _, err := service.store.GetClientView(ctx, request.GetClientId()); err != nil { + return nil, mapStoreError(err) + } + limit, err := pageSize(request.GetPageSize()) + if err != nil { + return nil, err + } + filterBytes := []byte(fmt.Sprintf("commands\x00%s\x00%t", request.GetClientId(), request.GetIncludeTerminal())) + filter := domain.HashCursorFilters(filterBytes) + page := store.CommandPage{ClientID: request.GetClientId(), IncludeTerminal: request.GetIncludeTerminal(), Limit: limit, SnapshotBoundary: service.now().UnixNano()} + if request.GetPageToken() != "" { + cursor, decodeErr := service.cursors.Decode(request.GetPageToken(), domain.CursorKindCommands, filter) + if decodeErr != nil || len(cursor.Position) != 24 || len(cursor.SnapshotBoundary) != 8 { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "invalid command page token") + } + page.SnapshotBoundary = int64(binary.BigEndian.Uint64(cursor.SnapshotBoundary)) + page.AfterTime = int64(binary.BigEndian.Uint64(cursor.Position[:8])) + copy(page.AfterUUID[:], cursor.Position[8:]) + page.HasAfter = true + } + commands, hasNext, err := service.store.ListCommandViews(ctx, page) + if err != nil { + return nil, mapStoreError(err) + } + response := &rvboxv1.ListCommandsResponse{Commands: make([]*rvboxv1.CommandRecord, 0, len(commands))} + for _, command := range commands { + converted, convertErr := commandRecord(command) + if convertErr != nil { + return nil, mapStoreError(convertErr) + } + response.Commands = append(response.Commands, converted) + } + if hasNext && len(commands) > 0 { + position := make([]byte, 24) + binary.BigEndian.PutUint64(position[:8], uint64(commands[len(commands)-1].IssueTime.UnixNano())) + copy(position[8:], commands[len(commands)-1].IssueUUID[:]) + boundary := make([]byte, 8) + binary.BigEndian.PutUint64(boundary, uint64(page.SnapshotBoundary)) + response.NextPageToken, err = service.cursors.Encode(domain.Cursor{Kind: domain.CursorKindCommands, FilterHash: filter, Position: position, SnapshotBoundary: boundary}) + if err != nil { + return nil, mapStoreError(err) + } + } + return response, nil +} + +func (service *Service) GetCommand(ctx context.Context, request *rvboxv1.GetCommandRequest) (*rvboxv1.GetCommandResponse, error) { + if request == nil || request.GetIssueUuid() == "" { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "issue_uuid is required") + } + issue, err := domain.ParseUUIDv7(request.GetIssueUuid()) + if err != nil { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "issue_uuid must be canonical UUIDv7") + } + command, err := service.store.GetCommandView(ctx, request.GetClientId(), issue) + if err != nil { + return nil, mapStoreError(err) + } + record, err := commandRecord(command) + if err != nil { + return nil, mapStoreError(err) + } + return &rvboxv1.GetCommandResponse{Command: record}, nil +} + +// RunCommand durably admits command-text work. Script payload persistence and +// chunk dispatch are intentionally kept behind the same immutable boundary and +// are rejected until their command_payloads path is wired to the dispatcher. +func (service *Service) RunCommand(ctx context.Context, request *rvboxv1.RunCommandRequest) (*rvboxv1.RunCommandResponse, error) { + if request == nil || request.GetTargetClientId() == "" || request.GetSpec() == nil { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "target_client_id and spec are required") + } + client, err := service.store.GetClientView(ctx, request.GetTargetClientId()) + if err != nil { + return nil, mapStoreError(err) + } + platform := rvboxv1.Platform(client.Platform) + spec := proto.Clone(request.GetSpec()).(*rvboxv1.ExecutionSpec) + if spec.GetShellType() == rvboxv1.ShellType_SHELL_TYPE_UNSPECIFIED { + switch platform { + case rvboxv1.Platform_PLATFORM_WINDOWS: + spec.ShellType = rvboxv1.ShellType_SHELL_POWERSHELL + case rvboxv1.Platform_PLATFORM_LINUX, rvboxv1.Platform_PLATFORM_DARWIN, rvboxv1.Platform_PLATFORM_OTHER_UNIX: + spec.ShellType = rvboxv1.ShellType_SHELL_SH + default: + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "target client platform is unspecified") + } + } + if _, script := spec.Source.(*rvboxv1.ExecutionSpec_Script); script { + return nil, status.Error(codes.Unimplemented, "script command admission is not enabled until payload dispatch is implemented") + } + if len(request.GetScriptContent()) != 0 { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "script_content requires a script source") + } + if err := agentproto.ValidateExecutionSpec(spec, service.limits, platform); err != nil { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, err.Error()) + } + if !advertisedShell(client.SupportedShells, spec.GetShellType()) { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "target client does not advertise the requested shell") + } + issue, err := service.requestID(request.GetRequestId()) + if err != nil { + return nil, err + } + now := service.now().UTC() + expiry, err := service.queueExpiry(request.GetQueueTtl(), now) + if err != nil { + return nil, err + } + canonical := &rvboxv1.RunCommandRequest{TargetClientId: request.GetTargetClientId(), Spec: spec, QueueTtl: request.GetQueueTtl()} + encoded, err := proto.MarshalOptions{Deterministic: true}.Marshal(canonical) + if err != nil { + return nil, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "canonicalize command request") + } + hash := sha256.Sum256(encoded) + queued, err := service.store.QueueCommand(ctx, store.QueueCommandInput{ + IssueUUID: issue, ClientID: request.GetTargetClientId(), IssueTime: now, ReceiptTime: now, + QueueExpiryTime: expiry, ImmutableSHA256: hash, ExecutionSpec: mustMarshal(spec), + }) + if err != nil { + return nil, mapStoreError(err) + } + return &rvboxv1.RunCommandResponse{IssueUuid: issue.String(), Lifecycle: rvboxv1.CommandLifecycle(queued.Lifecycle)}, nil +} + +func (service *Service) requestID(value string) (domain.UUID, error) { + if value == "" { + issue, err := domain.NewUUIDv7() + if err != nil { + return domain.UUID{}, controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, "generate request ID") + } + return issue, nil + } + issue, err := domain.ParseUUIDv7(value) + if err != nil { + return domain.UUID{}, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "request_id must be canonical UUIDv7") + } + return issue, nil +} + +func (service *Service) queueExpiry(value *durationpb.Duration, now time.Time) (*time.Time, error) { + if value == nil { + expiry := now.Add(service.defaultQueueTTL) + return &expiry, nil + } + if err := value.CheckValid(); err != nil || value.AsDuration() < 0 { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "queue_ttl must be non-negative") + } + if value.AsDuration() == 0 { + return nil, nil + } + duration := value.AsDuration() + if now.UnixNano() > math.MaxInt64-int64(duration) { + return nil, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "queue_ttl overflows server clock") + } + expiry := now.Add(duration) + return &expiry, nil +} + +func mustMarshal(message proto.Message) []byte { + encoded, _ := (proto.MarshalOptions{Deterministic: true}).Marshal(message) + return encoded +} + +func advertisedShell(encoded []byte, shell rvboxv1.ShellType) bool { + var hello rvboxv1.ClientHello + if err := proto.Unmarshal(encoded, &hello); err == nil { + for _, advertised := range hello.GetSupportedShells() { + if advertised == shell { + return true + } + } + } + return false +} + +func clientSummary(view store.ClientView) (*rvboxv1.ClientSummary, error) { + result := &rvboxv1.ClientSummary{ClientId: view.ClientID, Connected: view.Connected, Platform: rvboxv1.Platform(view.Platform), Architecture: view.Architecture, DaemonVersion: view.DaemonVersion, DaemonCwd: view.DaemonCWD, RunningCommands: view.RunningCommands, QueuedCommands: view.QueuedCommands, ClientInstanceId: uuidString(view.ClientInstanceID)} + var hello rvboxv1.ClientHello + if err := proto.Unmarshal(view.SupportedShells, &hello); err != nil { + return nil, fmt.Errorf("decode stored shell advertisement: %w", err) + } + result.SupportedShells = append([]rvboxv1.ShellType(nil), hello.GetSupportedShells()...) + if view.ConnectedAt != nil { + result.ConnectedAt = timestamppb.New(*view.ConnectedAt) + } + if view.LastSeenAt != nil { + result.LastSeenAt = timestamppb.New(*view.LastSeenAt) + } + if view.PendingInstanceID != nil { + result.PendingInstanceId = uuidString(*view.PendingInstanceID) + } + if view.PendingInstanceSeenAt != nil { + result.PendingInstanceSeenAt = timestamppb.New(*view.PendingInstanceSeenAt) + } + return result, nil +} + +func commandRecord(view store.CommandView) (*rvboxv1.CommandRecord, error) { + spec := &rvboxv1.ExecutionSpec{} + if err := proto.Unmarshal(view.ExecutionSpec, spec); err != nil { + return nil, fmt.Errorf("decode stored execution spec: %w", err) + } + result := &rvboxv1.CommandRecord{IssueUuid: view.IssueUUID.String(), TargetClientId: view.ClientID, IssueTime: timestamppb.New(view.IssueTime), ServerReceiptTime: timestamppb.New(view.ServerReceiptTime), Spec: spec, Lifecycle: rvboxv1.CommandLifecycle(view.Lifecycle), LastEventSeq: view.LastEventSeq, OutputTruncated: view.OutputTruncated, OutputIncomplete: view.OutputIncomplete, RetainedCompressedBytes: view.RetainedCompressedBytes, CommandRevision: view.Revision} + if view.QueueExpiryTime != nil { + result.QueueExpiryTime = timestamppb.New(*view.QueueExpiryTime) + } + if view.TerminalTime != nil { + result.TerminalTime = timestamppb.New(*view.TerminalTime) + } + if view.ExitCode != nil { + result.ExitCode = view.ExitCode + } + if len(view.WindowsIdentity) > 0 { + identity := &rvboxv1.WindowsExecutionIdentity{} + if err := proto.Unmarshal(view.WindowsIdentity, identity); err != nil { + return nil, fmt.Errorf("decode stored Windows identity: %w", err) + } + result.WindowsExecutionIdentity = identity + } + return result, nil +} + +func uuidString(value [16]byte) string { return domain.UUID(value).String() } + +func pageSize(value uint32) (uint32, error) { + if value == 0 { + return defaultPageSize, nil + } + if value > maxPageSize { + return 0, controlError(codes.InvalidArgument, rvboxv1.ControlError_INVALID_ARGUMENT, "page_size exceeds the maximum") + } + return value, nil +} + +func controlError(code codes.Code, rvCode rvboxv1.ControlError_Code, message string) error { + detail := &rvboxv1.ControlError{Code: rvCode, Message: message, Retryable: rvCode == rvboxv1.ControlError_OFFLINE || rvCode == rvboxv1.ControlError_CAPACITY_EXHAUSTED || rvCode == rvboxv1.ControlError_TRANSIENT || rvCode == rvboxv1.ControlError_INTERNAL} + result := status.New(code, message) + withDetails, err := result.WithDetails(detail) + if err == nil { + return withDetails.Err() + } + return result.Err() +} + +func mapStoreError(err error) error { + if err == nil { + return nil + } + switch { + case errors.Is(err, store.ErrClientNotFound), errors.Is(err, store.ErrCommandNotFound), errors.Is(err, store.ErrIncidentNotFound): + return controlError(codes.NotFound, rvboxv1.ControlError_NOT_FOUND, err.Error()) + case errors.Is(err, store.ErrCommandConflict), errors.Is(err, store.ErrMutationConflict): + return controlError(codes.AlreadyExists, rvboxv1.ControlError_CONFLICT, err.Error()) + case errors.Is(err, store.ErrCapacityExhausted): + return controlError(codes.ResourceExhausted, rvboxv1.ControlError_CAPACITY_EXHAUSTED, err.Error()) + case errors.Is(err, store.ErrStoreClosed): + return controlError(codes.Unavailable, rvboxv1.ControlError_OFFLINE, err.Error()) + default: + return controlError(codes.Internal, rvboxv1.ControlError_INTERNAL, err.Error()) + } +} diff --git a/internal/server/control/service_test.go b/internal/server/control/service_test.go new file mode 100644 index 0000000..479d496 --- /dev/null +++ b/internal/server/control/service_test.go @@ -0,0 +1,181 @@ +package control + +import ( + "context" + "crypto/sha256" + "net" + "path/filepath" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "github.com/rvbox/rvbox/internal/domain" + "github.com/rvbox/rvbox/internal/server/store" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/status" + "google.golang.org/grpc/test/bufconn" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/durationpb" +) + +func TestControlListAndGetViews_HP_CONTROL_01(t *testing.T) { + service, persistence := newTestService(t) + defer persistence.Close() + ctx := context.Background() + registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_POWERSHELL, 1) + registerControlClient(t, persistence, "win-b", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_CMD, 2) + + first, err := service.ListClients(ctx, &rvboxv1.ListClientsRequest{PageSize: 1}) + if err != nil || len(first.GetClients()) != 1 || first.GetNextPageToken() == "" { + t.Fatalf("first client page = %#v, %v", first, err) + } + second, err := service.ListClients(ctx, &rvboxv1.ListClientsRequest{PageSize: 1, PageToken: first.GetNextPageToken()}) + if err != nil || len(second.GetClients()) != 1 || second.GetClients()[0].GetClientId() != "win-b" { + t.Fatalf("second client page = %#v, %v", second, err) + } + if _, err := service.ListClients(ctx, &rvboxv1.ListClientsRequest{PageSize: 1, PageToken: "bad-token"}); status.Code(err) != codes.InvalidArgument { + t.Fatalf("bad client cursor code = %v", status.Code(err)) + } + client, err := service.GetClient(ctx, &rvboxv1.GetClientRequest{ClientId: "win-a"}) + if err != nil || client.GetClient().GetSupportedShells()[0] != rvboxv1.ShellType_SHELL_POWERSHELL { + t.Fatalf("get client = %#v, %v", client, err) + } + + issue := fixedIssue(0xa1) + queued, err := service.RunCommand(ctx, &rvboxv1.RunCommandRequest{ + TargetClientId: "win-a", RequestId: issue.String(), + Spec: &rvboxv1.ExecutionSpec{Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo hi"}}, + }) + if err != nil || queued.GetIssueUuid() != issue.String() || queued.GetLifecycle() != rvboxv1.CommandLifecycle_COMMAND_QUEUED { + t.Fatalf("run response = %#v, %v", queued, err) + } + commands, err := service.ListCommands(ctx, &rvboxv1.ListCommandsRequest{ClientId: "win-a", IncludeTerminal: true}) + if err != nil || len(commands.GetCommands()) != 1 || commands.GetCommands()[0].GetSpec().GetShellType() != rvboxv1.ShellType_SHELL_POWERSHELL { + t.Fatalf("command list = %#v, %v", commands, err) + } + got, err := service.GetCommand(ctx, &rvboxv1.GetCommandRequest{ClientId: "win-a", IssueUuid: issue.String()}) + if err != nil || got.GetCommand().GetIssueUuid() != issue.String() { + t.Fatalf("command get = %#v, %v", got, err) + } + if _, err := service.GetCommand(ctx, &rvboxv1.GetCommandRequest{IssueUuid: "019c46f1-1d02-6000-8000-0000000000a1"}); status.Code(err) != codes.InvalidArgument { + t.Fatalf("non-v7 command ID code = %v", status.Code(err)) + } +} + +func TestRunCommandRequestIDIdempotencyAndValidation_BH_CONTROL_02(t *testing.T) { + service, persistence := newTestService(t) + defer persistence.Close() + ctx := context.Background() + registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_POWERSHELL, 3) + issue := fixedIssue(0xa2) + request := &rvboxv1.RunCommandRequest{TargetClientId: "win-a", RequestId: issue.String(), Spec: &rvboxv1.ExecutionSpec{Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo one"}}} + if _, err := service.RunCommand(ctx, request); err != nil { + t.Fatal(err) + } + if replay, err := service.RunCommand(ctx, proto.Clone(request).(*rvboxv1.RunCommandRequest)); err != nil || replay.GetIssueUuid() != issue.String() { + t.Fatalf("exact replay = %#v, %v", replay, err) + } + conflict := proto.Clone(request).(*rvboxv1.RunCommandRequest) + conflict.Spec = &rvboxv1.ExecutionSpec{Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo two"}} + if _, err := service.RunCommand(ctx, conflict); status.Code(err) != codes.AlreadyExists { + t.Fatalf("conflicting replay code = %v", status.Code(err)) + } + badID := proto.Clone(request).(*rvboxv1.RunCommandRequest) + badID.RequestId = "not-a-uuid" + if _, err := service.RunCommand(ctx, badID); status.Code(err) != codes.InvalidArgument { + t.Fatalf("bad request ID code = %v", status.Code(err)) + } + script := &rvboxv1.RunCommandRequest{TargetClientId: "win-a", RequestId: fixedIssue(0xa3).String(), Spec: &rvboxv1.ExecutionSpec{Source: &rvboxv1.ExecutionSpec_Script{Script: &rvboxv1.ScriptDescriptor{Filename: "x.ps1", SizeBytes: 1, Sha256: sha256.New().Sum(nil)}}}, ScriptContent: []byte("x")} + if _, err := service.RunCommand(ctx, script); status.Code(err) != codes.Unimplemented { + t.Fatalf("script admission code = %v", status.Code(err)) + } + badTTL := proto.Clone(request).(*rvboxv1.RunCommandRequest) + badTTL.RequestId = fixedIssue(0xa4).String() + badTTL.QueueTtl = durationpb.New(-time.Second) + if _, err := service.RunCommand(ctx, badTTL); status.Code(err) != codes.InvalidArgument { + t.Fatalf("negative TTL code = %v", status.Code(err)) + } +} + +func TestCommandPaginationBindsFilters_BH_CONTROL_03(t *testing.T) { + service, persistence := newTestService(t) + defer persistence.Close() + ctx := context.Background() + registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_CMD, 4) + for index := byte(0xb0); index < 0xb3; index++ { + if _, err := service.RunCommand(ctx, &rvboxv1.RunCommandRequest{TargetClientId: "win-a", RequestId: fixedIssue(index).String(), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_CMD, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo"}}}); err != nil { + t.Fatal(err) + } + } + first, err := service.ListCommands(ctx, &rvboxv1.ListCommandsRequest{ClientId: "win-a", PageSize: 1}) + if err != nil || len(first.GetCommands()) != 1 || first.GetNextPageToken() == "" { + t.Fatalf("first command page = %#v, %v", first, err) + } + changedFilter := &rvboxv1.ListCommandsRequest{ClientId: "win-a", IncludeTerminal: true, PageSize: 1, PageToken: first.GetNextPageToken()} + if _, err := service.ListCommands(ctx, changedFilter); status.Code(err) != codes.InvalidArgument { + t.Fatalf("changed-filter cursor code = %v", status.Code(err)) + } + second, err := service.ListCommands(ctx, &rvboxv1.ListCommandsRequest{ClientId: "win-a", PageSize: 1, PageToken: first.GetNextPageToken()}) + if err != nil || len(second.GetCommands()) != 1 || second.GetCommands()[0].GetIssueUuid() == first.GetCommands()[0].GetIssueUuid() { + t.Fatalf("second command page = %#v, %v", second, err) + } +} + +func TestControlGRPCRoundTrip_HP_CONTROL_06(t *testing.T) { + service, persistence := newTestService(t) + defer persistence.Close() + registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_CMD, 5) + grpcServer := grpc.NewServer() + rvboxv1.RegisterControlServer(grpcServer, service) + listener := bufconn.Listen(1 << 20) + go func() { _ = grpcServer.Serve(listener) }() + defer grpcServer.Stop() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + connection, err := grpc.DialContext(ctx, "bufnet", grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) { return listener.Dial() }), grpc.WithTransportCredentials(insecure.NewCredentials()), grpc.WithBlock()) + if err != nil { + t.Fatal(err) + } + defer connection.Close() + client := rvboxv1.NewControlClient(connection) + response, err := client.RunCommand(ctx, &rvboxv1.RunCommandRequest{TargetClientId: "win-a", RequestId: fixedIssue(0xa5).String(), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_CMD, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo wire"}}}) + if err != nil || response.GetLifecycle() != rvboxv1.CommandLifecycle_COMMAND_QUEUED { + t.Fatalf("gRPC run response = %#v, %v", response, err) + } +} + +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}) + if err != nil { + t.Fatal(err) + } + service, err := NewService(Options{Store: persistence, CursorKey: bytesKey(), Now: func() time.Time { return time.Date(2026, time.September, 6, 12, 0, 0, 0, time.UTC) }, Limits: agentproto.DefaultLimits()}) + if err != nil { + persistence.Close() + t.Fatal(err) + } + return service, persistence +} + +func bytesKey() []byte { return []byte("0123456789abcdef0123456789abcdef") } + +func registerControlClient(t *testing.T, persistence *store.Store, clientID string, platform rvboxv1.Platform, shell rvboxv1.ShellType, seed byte) { + t.Helper() + shells, err := proto.Marshal(&rvboxv1.ClientHello{SupportedShells: []rvboxv1.ShellType{shell}}) + if err != nil { + t.Fatal(err) + } + if _, err := persistence.RegisterClientSession(context.Background(), store.ClientRegistration{ClientID: clientID, Platform: uint32(platform), Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\\ProgramData\\RVBox`, SupportedShells: shells, ClientInstanceID: [16]byte{seed}, SessionID: [16]byte{seed + 10}, ConnectedAt: time.Date(2026, time.September, 6, 11, 0, 0, 0, time.UTC)}); err != nil { + t.Fatal(err) + } +} + +func fixedIssue(last byte) domain.UUID { + value, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000001") + value[15] = last + return value +} diff --git a/internal/server/store/query.go b/internal/server/store/query.go new file mode 100644 index 0000000..ea4ad03 --- /dev/null +++ b/internal/server/store/query.go @@ -0,0 +1,284 @@ +package store + +import ( + "context" + "database/sql" + "errors" + "time" + + "github.com/rvbox/rvbox/internal/domain" +) + +// ClientView is the read-only, protocol-neutral representation of a client. +// The store deliberately returns copies of all byte slices so callers cannot +// mutate memory owned by a database driver or a shared scan buffer. +type ClientView struct { + ClientID string + Connected bool + ConnectedAt *time.Time + LastSeenAt *time.Time + RunningCommands uint32 + QueuedCommands uint32 + Platform uint32 + Architecture string + DaemonVersion string + DaemonCWD string + SupportedShells []byte + ClientInstanceID [16]byte + PendingInstanceID *[16]byte + PendingInstanceSeenAt *time.Time +} + +// CommandView is the durable command metadata exposed to the control layer. +// ExecutionSpec is returned in decoded form; the compressed representation +// never crosses the store boundary. +type CommandView struct { + IssueUUID domain.UUID + ClientID string + IssueTime time.Time + ServerReceiptTime time.Time + QueueExpiryTime *time.Time + TerminalTime *time.Time + Lifecycle uint32 + Revision uint64 + LastEventSeq uint64 + ExitCode *int32 + OutputTruncated bool + OutputIncomplete bool + RetainedCompressedBytes uint64 + ExecutionSpec []byte + WindowsIdentity []byte +} + +// ListClientViews returns clients ordered by client_id. The afterClientID +// value is an exclusive lexical cursor; an empty value starts at the first +// row. The boolean reports whether another row exists. +func (store *Store) ListClientViews(ctx context.Context, afterClientID string, limit uint32) ([]ClientView, bool, error) { + if limit == 0 || limit > 1000 { + return nil, false, errors.New("invalid client page size") + } + database, err := store.openDatabase() + if err != nil { + return nil, false, err + } + rows, err := database.QueryContext(ctx, `SELECT + c.client_id, c.platform, c.architecture, c.daemon_version, c.daemon_cwd, + c.supported_shells, c.client_instance_id, c.connected_at, c.last_seen_at, + c.pending_instance_id, c.pending_instance_seen_at, + EXISTS(SELECT 1 FROM sessions s WHERE s.client_id = c.client_id AND s.closed_at IS NULL AND s.fenced_at IS NULL), + (SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (3,4)), + (SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (1,2)) + FROM clients c WHERE c.client_id > ? ORDER BY c.client_id LIMIT ?`, afterClientID, limit+1) + if err != nil { + return nil, false, err + } + defer rows.Close() + clients := make([]ClientView, 0, limit) + for rows.Next() { + view, err := scanClientView(rows) + if err != nil { + return nil, false, err + } + if uint32(len(clients)) < limit { + clients = append(clients, view) + } else { + return clients, true, rows.Err() + } + } + if err := rows.Err(); err != nil { + return nil, false, err + } + return clients, false, nil +} + +// GetClientView returns one client or ErrClientNotFound. Connection state is +// derived from the live-session row rather than the historical timestamps. +func (store *Store) GetClientView(ctx context.Context, clientID string) (ClientView, error) { + if clientID == "" { + return ClientView{}, ErrClientNotFound + } + database, err := store.openDatabase() + if err != nil { + return ClientView{}, err + } + row := database.QueryRowContext(ctx, `SELECT + c.client_id, c.platform, c.architecture, c.daemon_version, c.daemon_cwd, + c.supported_shells, c.client_instance_id, c.connected_at, c.last_seen_at, + c.pending_instance_id, c.pending_instance_seen_at, + EXISTS(SELECT 1 FROM sessions s WHERE s.client_id = c.client_id AND s.closed_at IS NULL AND s.fenced_at IS NULL), + (SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (3,4)), + (SELECT COUNT(*) FROM commands q WHERE q.client_id = c.client_id AND q.lifecycle IN (1,2)) + FROM clients c WHERE c.client_id = ?`, clientID) + view, err := scanClientView(row) + if errors.Is(err, sql.ErrNoRows) { + return ClientView{}, ErrClientNotFound + } + return view, err +} + +func scanClientView(scanner interface{ Scan(...any) error }) (ClientView, error) { + var view ClientView + var instance, pending []byte + var connectedAt, lastSeen, pendingSeen sql.NullInt64 + var connected, running, queued int64 + if err := scanner.Scan(&view.ClientID, &view.Platform, &view.Architecture, &view.DaemonVersion, &view.DaemonCWD, + &view.SupportedShells, &instance, &connectedAt, &lastSeen, &pending, &pendingSeen, + &connected, &running, &queued); err != nil { + return view, err + } + if len(instance) != 16 || (pending != nil && len(pending) != 16) { + return view, ErrInvalidSegmentRecord + } + copy(view.ClientInstanceID[:], instance) + if pending != nil { + var value [16]byte + copy(value[:], pending) + view.PendingInstanceID = &value + } + view.Connected = connected == 1 + if running < 0 || running > int64(^uint32(0)) || queued < 0 || queued > int64(^uint32(0)) { + return view, ErrInvalidSegmentRecord + } + view.RunningCommands, view.QueuedCommands = uint32(running), uint32(queued) + view.ConnectedAt = nullableTime(connectedAt) + view.LastSeenAt = nullableTime(lastSeen) + view.PendingInstanceSeenAt = nullableTime(pendingSeen) + view.SupportedShells = append([]byte(nil), view.SupportedShells...) + return view, nil +} + +func nullableTime(value sql.NullInt64) *time.Time { + if !value.Valid { + return nil + } + instant := time.Unix(0, value.Int64).UTC() + return &instant +} + +// CommandPage describes a stable descending command cursor. SnapshotBoundary +// is an inclusive issue-time ceiling captured by the first page. AfterTime and +// AfterUUID are the exclusive position from the previous page. +type CommandPage struct { + ClientID string + IncludeTerminal bool + Limit uint32 + SnapshotBoundary int64 + AfterTime int64 + AfterUUID domain.UUID + HasAfter bool +} + +// ListCommandViews reads a stable, descending command page. The caller must +// preserve SnapshotBoundary and the returned final (issue_time, UUID) pair in +// its authenticated cursor. +func (store *Store) ListCommandViews(ctx context.Context, page CommandPage) ([]CommandView, bool, error) { + if page.ClientID == "" || page.Limit == 0 || page.Limit > 1000 || page.SnapshotBoundary <= 0 { + return nil, false, errors.New("invalid command page") + } + database, err := store.openDatabase() + if err != nil { + return nil, false, err + } + query := `SELECT issue_uuid, client_id, issue_time, server_receipt_time, + queue_expiry_time, terminal_time, lifecycle, revision, last_event_seq, + exit_code, output_truncated, output_incomplete, retained_compressed_bytes, + execution_spec, execution_spec_raw_bytes, windows_execution_identity + FROM commands WHERE client_id = ? AND issue_time <= ?` + args := []any{page.ClientID, page.SnapshotBoundary} + if !page.IncludeTerminal { + query += ` AND lifecycle NOT BETWEEN 5 AND 11` + } + if page.HasAfter { + query += ` AND (issue_time < ? OR (issue_time = ? AND issue_uuid < ?))` + args = append(args, page.AfterTime, page.AfterTime, page.AfterUUID[:]) + } + query += ` ORDER BY issue_time DESC, issue_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() + commands := make([]CommandView, 0, page.Limit) + for rows.Next() { + view, err := scanCommandView(rows) + if err != nil { + return nil, false, err + } + if uint32(len(commands)) < page.Limit { + commands = append(commands, view) + } else { + return commands, true, rows.Err() + } + } + if err := rows.Err(); err != nil { + return nil, false, err + } + return commands, false, nil +} + +// GetCommandView returns one retained command, optionally constrained to its +// client. Evicted commands are represented by ErrCommandNotFound; the +// tombstone remains available to reconciliation rather than control history. +func (store *Store) GetCommandView(ctx context.Context, clientID string, issue domain.UUID) (CommandView, error) { + if issue == (domain.UUID{}) { + return CommandView{}, ErrCommandNotFound + } + database, err := store.openDatabase() + if err != nil { + return CommandView{}, err + } + query := `SELECT issue_uuid, client_id, issue_time, server_receipt_time, + queue_expiry_time, terminal_time, lifecycle, revision, last_event_seq, + exit_code, output_truncated, output_incomplete, retained_compressed_bytes, + execution_spec, execution_spec_raw_bytes, windows_execution_identity + FROM commands WHERE issue_uuid = ?` + args := []any{issue[:]} + if clientID != "" { + query += ` AND client_id = ?` + args = append(args, clientID) + } + view, err := scanCommandView(database.QueryRowContext(ctx, query, args...)) + if errors.Is(err, sql.ErrNoRows) { + return CommandView{}, ErrCommandNotFound + } + return view, err +} + +func scanCommandView(scanner interface{ Scan(...any) error }) (CommandView, error) { + var view CommandView + var issue, stored []byte + var issueTime, receipt int64 + var expiry, terminal sql.NullInt64 + var exit sql.NullInt64 + var truncated, incomplete int + var rawBytes uint64 + if err := scanner.Scan(&issue, &view.ClientID, &issueTime, &receipt, &expiry, &terminal, + &view.Lifecycle, &view.Revision, &view.LastEventSeq, &exit, &truncated, &incomplete, + &view.RetainedCompressedBytes, &stored, &rawBytes, &view.WindowsIdentity); err != nil { + return view, err + } + if len(issue) != 16 { + return view, ErrInvalidSegmentRecord + } + copy(view.IssueUUID[:], issue) + view.IssueTime = time.Unix(0, issueTime).UTC() + view.ServerReceiptTime = time.Unix(0, receipt).UTC() + view.QueueExpiryTime = nullableTime(expiry) + view.TerminalTime = nullableTime(terminal) + view.OutputTruncated, view.OutputIncomplete = truncated == 1, incomplete == 1 + if exit.Valid { + if exit.Int64 < -1<<31 || exit.Int64 > 1<<31-1 { + return view, ErrInvalidSegmentRecord + } + value := int32(exit.Int64) + view.ExitCode = &value + } + var err error + view.ExecutionSpec, err = decompressCommandSpec(stored, rawBytes) + if err != nil { + return view, err + } + view.WindowsIdentity = append([]byte(nil), view.WindowsIdentity...) + return view, nil +} diff --git a/test/coverage.toml b/test/coverage.toml index 0d80231..65e1498 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -23,6 +23,42 @@ layer = "unit" status = "implemented" tests = ["internal/domain/cursor_test.go:TestAuthenticatedCursorRoundTrip_HP_CTL_06"] +[[requirements]] +id = "HP-CTL-07" +layer = "unit" +status = "implemented" +tests = ["internal/server/control/service_test.go:TestControlListAndGetViews_HP_CONTROL_01"] + +[[requirements]] +id = "BH-CTL-02" +layer = "unit" +status = "implemented" +tests = ["internal/server/control/service_test.go:TestRunCommandRequestIDIdempotencyAndValidation_BH_CONTROL_02"] + +[[requirements]] +id = "BH-CTL-03" +layer = "unit" +status = "implemented" +tests = ["internal/server/control/service_test.go:TestCommandPaginationBindsFilters_BH_CONTROL_03"] + +[[requirements]] +id = "HP-CTL-08" +layer = "integration" +status = "implemented" +tests = ["internal/server/control/service_test.go:TestControlGRPCRoundTrip_HP_CONTROL_06"] + +[[requirements]] +id = "HP-CTL-09" +layer = "unit" +status = "implemented" +tests = ["internal/server/control/listen_unix_test.go:TestListenUnixProtectsAndCleansOwnedSocket_HP_CONTROL_04"] + +[[requirements]] +id = "BH-CTL-04" +layer = "unit" +status = "implemented" +tests = ["internal/server/control/listen_unix_test.go:TestListenUnixRefusesNonSocketPath_BH_CONTROL_05"] + [[requirements]] id = "BH-CTL-01" layer = "unit"