feat: add Unix control service and rvc queue path
This commit is contained in:
@@ -13,10 +13,13 @@ import (
|
|||||||
"os/signal"
|
"os/signal"
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
|
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
|
||||||
"github.com/rvbox/rvbox/internal/agentproto"
|
"github.com/rvbox/rvbox/internal/agentproto"
|
||||||
"github.com/rvbox/rvbox/internal/config"
|
"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/session"
|
||||||
"github.com/rvbox/rvbox/internal/server/store"
|
"github.com/rvbox/rvbox/internal/server/store"
|
||||||
|
"google.golang.org/grpc"
|
||||||
)
|
)
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
@@ -60,6 +63,21 @@ func run(configPath string) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("listen for agents: %w", err)
|
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{
|
agent := &session.AgentServer{
|
||||||
Store: persistence, Registry: session.NewRegistry(), Path: configured.Server.AgentPath,
|
Store: persistence, Registry: session.NewRegistry(), Path: configured.Server.AgentPath,
|
||||||
Limits: agentproto.Limits{
|
Limits: agentproto.Limits{
|
||||||
@@ -73,8 +91,9 @@ func run(configPath string) error {
|
|||||||
HeartbeatIdle: configured.Protocol.HeartbeatIdle, LivenessTimeout: configured.Protocol.LivenessTimeout,
|
HeartbeatIdle: configured.Protocol.HeartbeatIdle, LivenessTimeout: configured.Protocol.LivenessTimeout,
|
||||||
}
|
}
|
||||||
httpServer := &http.Server{Handler: agent, ReadHeaderTimeout: configured.Flow.WriteDeadline}
|
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 <- httpServer.Serve(listener) }()
|
||||||
|
go func() { serveError <- grpcServer.Serve(controlListener) }()
|
||||||
|
|
||||||
signals := make(chan os.Signal, 1)
|
signals := make(chan os.Signal, 1)
|
||||||
signal.Notify(signals, os.Interrupt, syscall.SIGTERM)
|
signal.Notify(signals, os.Interrupt, syscall.SIGTERM)
|
||||||
@@ -82,12 +101,22 @@ func run(configPath string) error {
|
|||||||
select {
|
select {
|
||||||
case err := <-serveError:
|
case err := <-serveError:
|
||||||
if errors.Is(err, http.ErrServerClosed) {
|
if errors.Is(err, http.ErrServerClosed) {
|
||||||
|
grpcServer.Stop()
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
grpcServer.Stop()
|
||||||
return err
|
return err
|
||||||
case <-signals:
|
case <-signals:
|
||||||
shutdownContext, cancel := context.WithTimeout(context.Background(), configured.Server.ShutdownGrace)
|
shutdownContext, cancel := context.WithTimeout(context.Background(), configured.Server.ShutdownGrace)
|
||||||
defer cancel()
|
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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+285
-2
@@ -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
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ const (
|
|||||||
CursorKindCommands CursorKind = iota + 1
|
CursorKindCommands CursorKind = iota + 1
|
||||||
CursorKindOutput
|
CursorKindOutput
|
||||||
CursorKindIncidents
|
CursorKindIncidents
|
||||||
|
CursorKindClients
|
||||||
)
|
)
|
||||||
|
|
||||||
type Cursor struct {
|
type Cursor struct {
|
||||||
@@ -108,5 +109,5 @@ func (codec *CursorCodec) Decode(token string, expectedKind CursorKind, expected
|
|||||||
}
|
}
|
||||||
|
|
||||||
func validCursorKind(kind CursorKind) bool {
|
func validCursorKind(kind CursorKind) bool {
|
||||||
return kind >= CursorKindCommands && kind <= CursorKindIncidents
|
return kind >= CursorKindCommands && kind <= CursorKindClients
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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())
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -23,6 +23,42 @@ layer = "unit"
|
|||||||
status = "implemented"
|
status = "implemented"
|
||||||
tests = ["internal/domain/cursor_test.go:TestAuthenticatedCursorRoundTrip_HP_CTL_06"]
|
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]]
|
[[requirements]]
|
||||||
id = "BH-CTL-01"
|
id = "BH-CTL-01"
|
||||||
layer = "unit"
|
layer = "unit"
|
||||||
|
|||||||
Reference in New Issue
Block a user