diff --git a/cmd/rvbox/main.go b/cmd/rvbox/main.go index 343ae55..c49c832 100644 --- a/cmd/rvbox/main.go +++ b/cmd/rvbox/main.go @@ -1,4 +1,222 @@ // Command rvbox is the RVBox client daemon and Windows service executable. package main -func main() {} +import ( + "context" + "crypto/tls" + "crypto/x509" + "errors" + "flag" + "fmt" + "io" + "net/http" + "os" + "runtime" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "github.com/rvbox/rvbox/internal/client/agent" + "github.com/rvbox/rvbox/internal/client/spool" + clientwindows "github.com/rvbox/rvbox/internal/client/supervisor/windows" + "github.com/rvbox/rvbox/internal/client/windowsservice" + "github.com/rvbox/rvbox/internal/config" + "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/types/known/timestamppb" +) + +func main() { + if err := run(os.Args[1:], os.Stdout, os.Stderr); err != nil { + fmt.Fprintln(os.Stderr, "rvbox:", err) + os.Exit(1) + } +} + +// run is deliberately a small mode dispatcher. The SCM service invokes only +// --service with an explicit config path; tray/helper modes cannot silently +// turn an ordinary process invocation into a privileged service. +func run(args []string, output, diagnostics io.Writer) error { + if len(args) == 0 { + return errors.New("an internal mode is required (use --help)") + } + if args[0] == "--help" || args[0] == "-h" { + _, err := io.WriteString(output, "usage: rvbox --service|--tray|--check-config|--install-service|--uninstall-service|--start-service|--stop-service --config PATH\n") + return err + } + flags := flag.NewFlagSet("rvbox", flag.ContinueOnError) + flags.SetOutput(io.Discard) + configPath := flags.String("config", defaultClientConfigPath(), "absolute client TOML configuration path") + serviceMode := flags.Bool("service", false, "run under the Windows Service Control Manager") + trayMode := flags.Bool("tray", false, "run the current user's notification-area frontend") + checkConfig := flags.Bool("check-config", false, "validate client configuration and exit") + install := flags.Bool("install-service", false, "install or update the machine-wide service") + uninstall := flags.Bool("uninstall-service", false, "remove the machine-wide service") + start := flags.Bool("start-service", false, "start the machine-wide service") + stop := flags.Bool("stop-service", false, "stop the machine-wide service") + if err := flags.Parse(args); err != nil { + return err + } + if flags.NArg() != 0 { + return fmt.Errorf("unexpected argument %q", flags.Arg(0)) + } + selected := 0 + for _, value := range []bool{*serviceMode, *trayMode, *checkConfig, *install, *uninstall, *start, *stop} { + if value { + selected++ + } + } + if selected != 1 { + return errors.New("select exactly one rvbox mode") + } + if *checkConfig { + if _, err := loadClientConfig(*configPath); err != nil { + return err + } + _, err := fmt.Fprintf(output, "valid client configuration: %s\n", *configPath) + return err + } + if *install { + executable, err := os.Executable() + if err != nil { + return err + } + return windowsservice.Install(windowsservice.InstallSpec{ExecutablePath: executable, ConfigPath: *configPath, Startup: windowsservice.StartupAutomatic}) + } + if *uninstall { + return windowsservice.Uninstall() + } + if *start { + return windowsservice.Start() + } + if *stop { + return windowsservice.Stop(30) + } + if *trayMode { + return runTray(*configPath, diagnostics) + } + return runService(*configPath, diagnostics) +} + +func loadClientConfig(path string) (*config.Client, error) { + platform := config.PlatformUnix + checkFilesystem := false + if runtime.GOOS == "windows" { + platform = config.PlatformWindows + checkFilesystem = true + } + data, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("read client config: %w", err) + } + return config.DecodeClient(data, config.ClientOptions{Platform: platform, CheckFilesystem: checkFilesystem}) +} + +func runClientDaemon(ctx context.Context, configPath string, diagnostics io.Writer) error { + configured, err := loadClientConfig(configPath) + if err != nil { + return err + } + state, err := spool.Open(ctx, spool.Options{DataDir: configured.Client.StateDir, BusyTimeout: 5 * time.Second, TombstoneLimit: configured.Storage.TombstoneMaxEntries, MaxScriptBytes: configured.Execution.MaxScriptBytes, MaxExecutionSpecBytes: configured.Execution.MaxExecutionSpecBytes, QuotaLimits: spool.QuotaLimits{HardAllocationBytes: 1 << 20, CommandOutputBytes: configured.Storage.CommandOutputLimitBytes, CommandTotalBytes: configured.Storage.CommandTotalLimitBytes, ClientTotalBytes: configured.Storage.ClientTotalLimitBytes, CloseoutReserveBytes: configured.Storage.CommandCloseoutReserveBytes}}) + if err != nil { + return fmt.Errorf("open client durable state: %w", err) + } + defer state.Close() + if diagnostics != nil { + _, _ = fmt.Fprintf(diagnostics, "rvbox client service initialized for %s\n", configured.Client.ServerURL) + } + hello := clientHello(configured, state.ClientInstanceID()) + httpClient, err := clientHTTPClient(configured.TLS) + if err != nil { + return fmt.Errorf("configure client TLS: %w", err) + } + limits := agentproto.Limits{MaxEnvelopeBytes: configured.Execution.MaxAgentEnvelopeBytes, MaxExecutionSpecBytes: configured.Execution.MaxExecutionSpecBytes, MaxRawChunkBytes: configured.Execution.MaxRawChunkBytes, MaxScriptBytes: configured.Execution.MaxScriptBytes, MaxDetailBytes: configured.Execution.ProtocolDetailMaxBytes} + if limits.MaxEnvelopeBytes == 0 { + limits = agentproto.DefaultLimits() + } + eventReady := make(chan domain.UUID, 256) + supervised, err := clientwindows.NewSupervisor(clientwindows.NativeOptions{Shells: clientwindows.ShellPaths{CMD: configured.Shells.CMD, PowerShell: configured.Shells.PowerShell}, WorkRoot: configured.Client.DaemonCWD, MaxWrapperBytes: configured.Execution.MaxScriptBytes, MaxOutputChunk: configured.Execution.MaxRawChunkBytes, WindowsTermGrace: configured.Execution.WindowsTermGrace}) + if err != nil { + return fmt.Errorf("configure command supervisor: %w", err) + } + executor, err := agent.NewExecutor(agent.ExecutorOptions{Store: state, Supervisor: supervised, WorkDir: configured.Client.DaemonCWD, Notify: func(issue domain.UUID) { + select { + case eventReady <- issue: + default: + } + }}) + if err != nil { + return fmt.Errorf("configure command executor: %w", err) + } + runner := func() { + if _, checkErr := state.Check(ctx); checkErr != nil { + if diagnostics != nil { + _, _ = fmt.Fprintf(diagnostics, "rvbox client spool is dirty: %v\n", checkErr) + } + return + } + if recovered, recoverErr := state.RecoverLaunchUncertainty(ctx, time.Now().UTC()); recoverErr != nil { + if diagnostics != nil { + _, _ = fmt.Fprintf(diagnostics, "rvbox launch recovery failed: %v\n", recoverErr) + } + return + } else if len(recovered) > 0 && diagnostics != nil { + _, _ = fmt.Fprintf(diagnostics, "rvbox recovered %d uncertain launch(es)\n", len(recovered)) + } + if runErr := agent.Run(ctx, agent.RunnerOptions{ + Store: state, + Dial: func(dialContext context.Context) (agent.Transport, error) { + return agent.DialWebSocket(dialContext, configured.Client.ServerURL, httpClient) + }, + Hello: hello, Limits: limits, + Backoff: agent.BackoffOptions{Initial: configured.Network.ReconnectInitial, Maximum: configured.Network.ReconnectMax, StableReset: configured.Network.StableSessionReset}, + Jitter: agent.CryptoJitter, Now: func() time.Time { return time.Now().UTC() }, + OnDispatch: executor.Dispatch, OnScriptReady: executor.ScriptReady, OnStdin: executor.Stdin, + OnCloseStdin: executor.CloseStdin, OnSignal: executor.Signal, OnTerminate: executor.Terminate, EventReady: eventReady, + }); runErr != nil && ctx.Err() == nil && diagnostics != nil { + _, _ = fmt.Fprintf(diagnostics, "rvbox client session stopped: %v\n", runErr) + } + } + go runner() + <-ctx.Done() + _ = executor.StopAll(context.Background()) + return nil +} + +func clientHTTPClient(settings config.TLS) (*http.Client, error) { + tlsConfig := &tls.Config{MinVersion: tls.VersionTLS12, ServerName: settings.ServerName} // #nosec G402 -- TLS 1.2 is the v1 floor. + if settings.CAFile != "" { + pem, err := os.ReadFile(settings.CAFile) + if err != nil { + return nil, err + } + pool, err := x509.SystemCertPool() + if err != nil || pool == nil { + pool = x509.NewCertPool() + } + if !pool.AppendCertsFromPEM(pem) { + return nil, errors.New("TLS CA file contains no certificates") + } + tlsConfig.RootCAs = pool + } + return &http.Client{Transport: &http.Transport{TLSClientConfig: tlsConfig}}, nil +} + +func clientHello(configured *config.Client, instance domain.UUID) *rvboxv1.ClientHello { + platform := rvboxv1.Platform_PLATFORM_LINUX + shells := []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_CMD, rvboxv1.ShellType_SHELL_POWERSHELL} + if runtime.GOOS == "windows" { + platform = rvboxv1.Platform_PLATFORM_WINDOWS + } else { + shells = []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_SH, rvboxv1.ShellType_SHELL_BASH} + } + clientID := configured.Client.ClientID + if clientID == "" { + clientID, _ = os.Hostname() + } + return &rvboxv1.ClientHello{ + ClientId: clientID, SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, + DaemonVersion: "v1", Platform: platform, Architecture: runtime.GOARCH, DaemonCwd: configured.Client.DaemonCWD, + SupportedShells: shells, ClientInstanceId: instance.String(), MaxRunningCommands: configured.Client.MaxRunningCommands, + MaxQueuedCommands: configured.Client.MaxQueuedCommands, SentAt: timestamppb.New(time.Now().UTC()), + } +} diff --git a/cmd/rvbox/main_test.go b/cmd/rvbox/main_test.go new file mode 100644 index 0000000..bb2244a --- /dev/null +++ b/cmd/rvbox/main_test.go @@ -0,0 +1,30 @@ +package main + +import ( + "bytes" + "testing" +) + +func TestClientModeSelectionRequiresExactlyOneMode_HP_WINCLI_01(t *testing.T) { + t.Parallel() + var output, diagnostics bytes.Buffer + if err := run([]string{"--help"}, &output, &diagnostics); err != nil || !bytes.Contains(output.Bytes(), []byte("--service")) { + t.Fatalf("help = %q, %v", output.String(), err) + } + if err := run(nil, &output, &diagnostics); err == nil { + t.Fatal("empty client invocation accepted") + } + if err := run([]string{"--service", "--check-config"}, &output, &diagnostics); err == nil { + t.Fatal("multiple client modes accepted") + } +} + +func TestNonWindowsServiceModesRemainExplicitlyUnsupported_BH_WINCLI_01(t *testing.T) { + t.Parallel() + if err := runService("", nil); err == nil { + t.Fatal("non-Windows service mode unexpectedly available") + } + if err := runTray("", nil); err == nil { + t.Fatal("non-Windows tray mode unexpectedly available") + } +} diff --git a/cmd/rvbox/service_other.go b/cmd/rvbox/service_other.go new file mode 100644 index 0000000..e7bf0e9 --- /dev/null +++ b/cmd/rvbox/service_other.go @@ -0,0 +1,18 @@ +//go:build !windows + +package main + +import ( + "errors" + "io" +) + +func defaultClientConfigPath() string { return "/etc/rvbox/client.toml" } + +func runService(string, io.Writer) error { + return errors.New("the v1 client service is implemented for Windows only") +} + +func runTray(string, io.Writer) error { + return errors.New("the v1 tray frontend is implemented for Windows only") +} diff --git a/cmd/rvbox/service_windows.go b/cmd/rvbox/service_windows.go new file mode 100644 index 0000000..99ca77f --- /dev/null +++ b/cmd/rvbox/service_windows.go @@ -0,0 +1,42 @@ +//go:build windows + +package main + +import ( + "context" + "errors" + "fmt" + "io" + "os" + "path/filepath" + + "github.com/rvbox/rvbox/internal/client/windowsservice" + "golang.org/x/sys/windows/svc" +) + +func defaultClientConfigPath() string { + root := os.Getenv("ProgramData") + if root == "" { + root = `C:\ProgramData` + } + return filepath.Join(root, "RVBox", "client.toml") +} + +func runService(configPath string, diagnostics io.Writer) error { + inService, err := svc.IsWindowsService() + if err != nil { + return fmt.Errorf("detect service control manager context: %w", err) + } + if !inService { + return errors.New("--service is reserved for the installed Windows service") + } + return runWindowsService(configPath, diagnostics) +} + +func runWindowsService(configPath string, diagnostics io.Writer) error { + return windowsservice.Run(func(ctx context.Context) error { return runClientDaemon(ctx, configPath, diagnostics) }) +} + +func runTray(configPath string, diagnostics io.Writer) error { + return errors.New("Windows tray frontend is not available in this build") +} diff --git a/docs/implementation-plan.v1.md b/docs/implementation-plan.v1.md index 6a5fd85..df5676b 100644 --- a/docs/implementation-plan.v1.md +++ b/docs/implementation-plan.v1.md @@ -1765,6 +1765,24 @@ rows in order. Advertise capacity and accept new dispatch only in `active`. Cancellation of one session context must join all its readers/writers before a new session can use their queues. +The current implementation checkpoint is intentionally split at this seam: +`internal/client/agent.RunOnce` owns the reconnect/session reader and remains +the sole live-session writer; `internal/client/spool` owns the SQLite source of +truth; and `internal/client/agent.Executor` owns process lifetime independently +of the WebSocket context. A bounded event-notification channel wakes the active +writer to assign and send newly appended events, while reconnect replay uses +the same spool rows and a per-session sent cursor. Every accepted dispatch is +stored with its deterministic execution specification (and, for scripts, its +descriptor reservation) before acknowledgement. + +Add a schema launch barrier to every implementation of the executor. Persist +`prepared` before entering the supervisor, persist `authorized` before the OS +release boundary, and clear it only when a terminal lifecycle transition is +committed. Startup recovery must convert any non-terminal `authorized` row to +one `interrupted` event before registering a new network session. This is the +at-most-once fence for a crash between process release and the first `running` +event; it is not a substitute for verifying a native process creation identity. + ### 7.2 Output capture and offline caps Create non-blocking readers for stdout and stderr immediately after process @@ -1917,6 +1935,15 @@ Implement Windows code in platform-specific files so non-Windows builds never import Windows APIs. Keep launch phases identical across platforms: `accepted -> launch_prepared -> launch_authorized -> running`, with no shortcut. +The first native adapter is now required to expose this contract through +`internal/client/supervisor.Supervisor`: the non-Windows adapter is test-only, +while the Windows implementation must perform token selection and Job setup +inside the same `Start` call. It may return only after the child has been +assigned to its kill-on-close Job and released; all token/session attempts must +be represented in the returned immutable identity. A failed start clears the +pre-launch barrier and produces one rejected lifecycle event; an uncertain +authorized row is never retried as a fresh process. + Implement one exhaustive token selector; do not scatter token fallback across launch code: diff --git a/docs/testing.md b/docs/testing.md index 5886742..e250de1 100644 --- a/docs/testing.md +++ b/docs/testing.md @@ -17,6 +17,8 @@ scripts/test-unit --package ./internal/domain --run UUIDv7 --race The integration harness provides the Phase 0 `sample` suite, the incremental Phase 2 `store` suite, and the incremental Phase 3 `server-session` suite. +The resumable E2E harness adds `smoke`, `script`, `recovery`, and `all` +scenarios. Each run writes its manifest and run ID before starting work. The storage suite uses a real temporary SQLite database in WAL mode and a real segment/audit filesystem. The session suite uses a real HTTP/WebSocket listener, binary protobuf frames, SQLite fencing, and the race detector; neither mocks its @@ -44,8 +46,25 @@ scripts/test-env logs --run-id session-smoke scripts/test-env collect --run-id session-smoke scripts/test-env reset --run-id session-smoke scripts/test-env purge --run-id session-smoke + +scripts/test-e2e --scenario smoke --run-id e2e-smoke +scripts/test-env status --run-id e2e-smoke +scripts/test-env recover --run-id e2e-smoke +scripts/test-e2e --scenario smoke --run-id e2e-smoke --resume +scripts/test-env reset --run-id e2e-smoke +scripts/test-env purge --run-id e2e-smoke ``` +The client runtime unit lane also exercises a real child process through the +portable supervisor adapter. `internal/client/agent/executor_test.go` verifies +that command text is accepted once, output is journaled, lifecycle/terminal +events are durable, and a script cannot launch before its contiguous upload is +committed. The Windows build uses the same executor contract with the +platform-native adapter: a verified token is selected, the child is created +suspended, assigned to a kill-on-close Job, and only then released. The +durable `launch_phase` barrier is recovered as `interrupted` after a daemon +restart, so an uncertain release is never redispatched. + Suite output is capped at 1 MiB and stored as `artifacts/suite.log`. A failed run remains inspectable and can be moved back to `ready` with `recover`, then resumed with the same run ID and deterministic shuffle seed. Test-run cleanup @@ -60,11 +79,12 @@ cleanup. `test/coverage.toml` is the incremental requirement-to-test inventory. The `make verify` lint stage checks unique stable IDs and verifies every implemented -test reference against source. A resettable Windows smoke VM is now available; -its exact headless VirtualBox/Guest Control runbook is in section 2.6.1 of -`docs/implementation-plan.v1.md`. Native Windows integration/E2E entries may -run there once the Windows harness acquires the exclusive lease and performs -the documented snapshot/health checks. The VM is only the minimum smoke lane, -so deferred native multi-session/ambiguous-session, Server Core, and -older-build entries remain explicitly blocked until their own fixtures exist. -Wine or a protocol stub is not treated as equivalent coverage. +test reference against source. A resettable Windows smoke VM is now available. +The exact headless VirtualBox/Guest Control adapter is +`scripts/windows/test-host.ps1`; it takes the VM name, baseline snapshot, and +guest credentials only from host environment variables, acquires an exclusive +lease, and never writes secrets to the repository. Use `Prepare`, `Run`, +`Collect`, `Stop`, and `Reset` in that order for a native run. The VM is the +minimum smoke lane, so deferred native multi-session/ambiguous-session, Server +Core, and older-build entries remain explicitly blocked until their own +fixtures exist. Wine or a protocol stub is not treated as equivalent coverage. diff --git a/internal/client/agent/dispatch.go b/internal/client/agent/dispatch.go index ee556fb..256a35f 100644 --- a/internal/client/agent/dispatch.go +++ b/internal/client/agent/dispatch.go @@ -35,15 +35,9 @@ func PersistDispatch(ctx context.Context, store *spool.Store, session Session, d if err != nil { return spool.Acceptance{}, err } - acceptance, err := store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: immutable, Revision: dispatch.GetCommandRevision(), Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED), ExecutionSpec: spec}, now) - if err != nil { - return spool.Acceptance{}, err - } + command := spool.Command{IssueUUID: issue, ImmutableSHA256: immutable, Revision: dispatch.GetCommandRevision(), Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED), ExecutionSpec: spec} if descriptor := dispatch.GetSpec().GetScript(); descriptor != nil { - _, err = store.BeginScript(ctx, issue, spool.ScriptDescriptor{SizeBytes: descriptor.GetSizeBytes(), SHA256: bytesToDigest(descriptor.GetSha256())}) - if err != nil { - return spool.Acceptance{}, err - } + command.Script = &spool.ScriptDescriptor{SizeBytes: descriptor.GetSizeBytes(), SHA256: bytesToDigest(descriptor.GetSha256())} } - return acceptance, nil + return store.AcceptCommand(ctx, command, now) } diff --git a/internal/client/agent/executor.go b/internal/client/agent/executor.go new file mode 100644 index 0000000..60539cf --- /dev/null +++ b/internal/client/agent/executor.go @@ -0,0 +1,375 @@ +package agent + +import ( + "context" + "errors" + "fmt" + "io" + "sync" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/client/spool" + "github.com/rvbox/rvbox/internal/client/supervisor" + "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/timestamppb" +) + +// Executor bridges durable dispatch records to the platform supervisor. It +// owns no network state: a command's process, pipes, and spool writes continue +// after the WebSocket session is replaced. +type Executor struct { + Store *spool.Store + Supervisor supervisor.Supervisor + WorkDir string + Now func() time.Time + Notify func(domain.UUID) + + mu sync.Mutex + active map[domain.UUID]supervisor.Process + cancel map[domain.UUID]context.CancelFunc + forced map[domain.UUID]bool +} + +type ExecutorOptions struct { + Store *spool.Store + Supervisor supervisor.Supervisor + WorkDir string + Now func() time.Time + Notify func(domain.UUID) +} + +func NewExecutor(options ExecutorOptions) (*Executor, error) { + if options.Store == nil || options.Supervisor == nil || options.WorkDir == "" { + return nil, errors.New("executor requires durable store, supervisor, and work directory") + } + if options.Now == nil { + options.Now = func() time.Time { return time.Now().UTC() } + } + return &Executor{Store: options.Store, Supervisor: options.Supervisor, WorkDir: options.WorkDir, Now: options.Now, Notify: options.Notify, active: make(map[domain.UUID]supervisor.Process), cancel: make(map[domain.UUID]context.CancelFunc), forced: make(map[domain.UUID]bool)}, nil +} + +// Dispatch is safe to invoke after CommandAccepted has been sent. Script +// commands intentionally wait for ScriptCommit; the server may deliver the +// descriptor and body in separate frames. +func (executor *Executor) Dispatch(ctx context.Context, _ Session, dispatch *rvboxv1.CommandDispatch) error { + if executor == nil || dispatch == nil { + return errors.New("invalid command dispatch") + } + issue, err := domain.ParseUUIDv7(dispatch.GetIssueUuid()) + if err != nil { + return err + } + if dispatch.GetSpec().GetScript() != nil { + if _, err := executor.Store.ScriptBody(ctx, issue); errors.Is(err, spool.ErrScriptNotReady) { + return nil + } else if err != nil { + return executor.reject(ctx, issue, dispatch.GetCommandRevision(), err) + } + } + return executor.launch(ctx, issue, dispatch.GetCommandRevision(), dispatch.GetSpec()) +} + +// ScriptReady launches a script after the durable commit barrier. Repeated +// progress/commit frames are harmless because the active map and spool phase +// make launch at-most-once. +func (executor *Executor) ScriptReady(ctx context.Context, _ Session, issue domain.UUID) error { + command, err := executor.Store.GetCommand(ctx, issue) + if err != nil { + return err + } + var spec rvboxv1.ExecutionSpec + if err := proto.Unmarshal(command.ExecutionSpec, &spec); err != nil { + return executor.reject(ctx, issue, command.Revision, err) + } + if spec.GetScript() == nil { + return nil + } + return executor.launch(ctx, issue, command.Revision, &spec) +} + +func (executor *Executor) launch(ctx context.Context, issue domain.UUID, revision uint64, spec *rvboxv1.ExecutionSpec) error { + if spec == nil || revision == 0 { + return executor.reject(ctx, issue, revision, errors.New("missing execution specification")) + } + executor.mu.Lock() + if _, exists := executor.active[issue]; exists { + executor.mu.Unlock() + return nil + } + executor.mu.Unlock() + var scriptBody []byte + if spec.GetScript() != nil { + body, err := executor.Store.ScriptBody(ctx, issue) + if err != nil { + if errors.Is(err, spool.ErrScriptNotReady) { + return nil + } + return executor.reject(ctx, issue, revision, err) + } + scriptBody = body + } + workingDirectory := spec.GetCwd() + if workingDirectory == "" { + workingDirectory = executor.WorkDir + } + if err := executor.Store.SetLaunchPhase(ctx, issue, domain.LaunchPhasePrepared, "", 0); err != nil { + return err + } + // The native Windows supervisor creates/assigns the Job while the child is + // suspended and releases it before returning. Marking authorization before + // that call makes a daemon crash in any of those windows recover as an + // interrupted, non-redispatchable command. + if err := executor.Store.SetLaunchPhase(ctx, issue, domain.LaunchPhaseAuthorized, "pending", 0); err != nil { + return err + } + runContext, cancel := context.WithCancel(context.Background()) + process, err := executor.Supervisor.Start(runContext, supervisor.StartSpec{IssueUUID: issue, CommandRevision: revision, Execution: proto.Clone(spec).(*rvboxv1.ExecutionSpec), ScriptBody: scriptBody, WorkingDirectory: workingDirectory, Environment: cloneEnvironment(spec.GetEnvOverrides()), ExecutionProfiles: executionProfileNames(spec.GetExecutionProfiles())}) + if err != nil { + cancel() + _ = executor.Store.SetLaunchPhase(context.Background(), issue, domain.LaunchPhaseNone, "", 0) + return executor.reject(ctx, issue, revision, err) + } + identity := process.Identity() + if err := executor.Store.SetLaunchPhase(ctx, issue, domain.LaunchPhaseAuthorized, identity.Context, 0); err != nil { + _, _ = executor.Supervisor.Signal(context.Background(), process, supervisor.SignalKill) + cancel() + return err + } + executor.mu.Lock() + executor.active[issue] = process + executor.cancel[issue] = cancel + executor.mu.Unlock() + if _, err := executor.Store.AppendLifecycle(ctx, issue, uint32(rvboxv1.CommandLifecycle_COMMAND_RUNNING), revision, "process started", executor.Now()); err != nil { + _, _ = executor.Supervisor.Signal(context.Background(), process, supervisor.SignalKill) + cancel() + executor.remove(issue) + return err + } + executor.notify(issue) + go executor.watch(runContext, issue, revision, process, cancel) + return nil +} + +func (executor *Executor) watch(ctx context.Context, issue domain.UUID, revision uint64, process supervisor.Process, cancel context.CancelFunc) { + defer cancel() + for { + chunk, err := process.ReadOutput(ctx) + if errors.Is(err, io.EOF) { + break + } + if err != nil { + _ = executor.appendIncomplete(context.Background(), issue, revision, err.Error()) + break + } + if len(chunk.Data) == 0 { + continue + } + if _, err := executor.Store.AppendOutput(context.Background(), issue, spool.OutputInput{Stream: chunk.Stream, Raw: chunk.Data, ObservedAt: executor.Now()}); err != nil { + _ = executor.appendIncomplete(context.Background(), issue, revision, err.Error()) + _, _ = executor.Supervisor.Signal(context.Background(), process, supervisor.SignalKill) + break + } + executor.notify(issue) + } + status, waitErr := process.Wait(context.Background()) + phase := rvboxv1.CommandLifecycle_COMMAND_FAILED + if waitErr == nil && status.Code == 0 { + phase = rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED + } + executor.mu.Lock() + if executor.forced[issue] { + phase = rvboxv1.CommandLifecycle_COMMAND_TERMINATED + } + executor.mu.Unlock() + detail := fmt.Sprintf("exit code %d", status.Code) + if waitErr != nil { + detail = boundedError(waitErr) + } + if err := executor.appendLifecycle(context.Background(), issue, revision, phase, detail); err == nil { + executor.notify(issue) + } + executor.remove(issue) +} + +func (executor *Executor) appendLifecycle(ctx context.Context, issue domain.UUID, revision uint64, phase rvboxv1.CommandLifecycle, detail string) error { + _, err := executor.Store.AppendLifecycle(ctx, issue, uint32(phase), revision, detail, executor.Now()) + return err +} + +func (executor *Executor) reject(ctx context.Context, issue domain.UUID, revision uint64, cause error) error { + if revision == 0 { + return cause + } + if err := executor.appendLifecycle(ctx, issue, revision, rvboxv1.CommandLifecycle_COMMAND_REJECTED, boundedError(cause)); err != nil { + return err + } + executor.notify(issue) + return nil +} + +func (executor *Executor) appendIncomplete(ctx context.Context, issue domain.UUID, _ uint64, detail string) error { + payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(&rvboxv1.CommandEvent{IssueUuid: issue.String(), ObservedAt: timestamppb.New(executor.Now()), Payload: &rvboxv1.CommandEvent_OutputIncomplete{OutputIncomplete: &rvboxv1.OutputIncomplete{Reason: boundedError(errors.New(detail))}}}) + if err != nil { + return err + } + _, err = executor.Store.AppendEvent(ctx, issue, spool.EventInput{Kind: 11, Compression: 1, RawBytes: uint64(len(payload)), Payload: payload, CreatedAt: executor.Now()}) + if err == nil { + executor.notify(issue) + } + return err +} + +func (executor *Executor) Stdin(ctx context.Context, _ Session, input *rvboxv1.StdinWrite) error { + if input == nil { + return errors.New("missing stdin request") + } + issue, err := domain.ParseUUIDv7(input.GetIssueUuid()) + if err != nil { + return err + } + executor.mu.Lock() + process := executor.active[issue] + executor.mu.Unlock() + if process == nil { + return executor.appendStdinAck(ctx, issue, input.GetWriteSeq(), false, "command is not running") + } + err = process.WriteStdin(ctx, input.GetData(), input.GetAppendNewline()) + detail := "" + if err != nil { + detail = boundedError(err) + } + return executor.appendStdinAck(ctx, issue, input.GetWriteSeq(), err == nil, detail) +} + +func (executor *Executor) CloseStdin(ctx context.Context, _ Session, input *rvboxv1.CloseStdin) error { + if input == nil { + return errors.New("missing close-stdin request") + } + issue, err := domain.ParseUUIDv7(input.GetIssueUuid()) + if err != nil { + return err + } + executor.mu.Lock() + process := executor.active[issue] + executor.mu.Unlock() + if process == nil { + return executor.appendStdinAck(ctx, issue, input.GetWriteSeq(), false, "command is not running") + } + err = process.CloseStdin(ctx) + detail := "" + if err != nil { + detail = boundedError(err) + } + return executor.appendStdinAck(ctx, issue, input.GetWriteSeq(), err == nil, detail) +} + +func (executor *Executor) appendStdinAck(ctx context.Context, issue domain.UUID, writeSeq uint64, accepted bool, detail string) error { + payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(&rvboxv1.CommandEvent{IssueUuid: issue.String(), ObservedAt: timestamppb.New(executor.Now()), Payload: &rvboxv1.CommandEvent_StdinAck{StdinAck: &rvboxv1.StdinAcknowledgement{WriteSeq: writeSeq, Detail: detail, StdinClosed: !accepted}}}) + if err != nil { + return err + } + _, err = executor.Store.AppendEvent(ctx, issue, spool.EventInput{Kind: 7, Compression: 1, RawBytes: uint64(len(payload)), Payload: payload, CreatedAt: executor.Now(), UseCloseout: false}) + if err == nil { + executor.notify(issue) + } + return err +} + +func (executor *Executor) Signal(ctx context.Context, _ Session, input *rvboxv1.SignalCommand) error { + if input == nil { + return errors.New("missing signal request") + } + issue, err := domain.ParseUUIDv7(input.GetIssueUuid()) + if err != nil { + return err + } + executor.mu.Lock() + process := executor.active[issue] + executor.mu.Unlock() + if process == nil { + return errors.New("command is not running") + } + kind := supervisor.SignalTerm + if input.GetSignal() == rvboxv1.SignalKind_SIGNAL_KILL { + kind = supervisor.SignalKill + } else if input.GetSignal() != rvboxv1.SignalKind_SIGNAL_TERM { + return errors.New("unsupported Windows signal") + } + outcome, err := executor.Supervisor.Signal(ctx, process, kind) + if kind == supervisor.SignalKill || outcome.Escalated { + executor.mu.Lock() + executor.forced[issue] = true + executor.mu.Unlock() + } + return executor.appendSignalResult(ctx, issue, input.GetCommandRevision(), input.GetSignal(), err == nil && outcome.Delivered, outcome, err) +} + +func (executor *Executor) appendSignalResult(ctx context.Context, issue domain.UUID, revision uint64, signal rvboxv1.SignalKind, accepted bool, outcome supervisor.SignalOutcome, cause error) error { + detail := outcome.Detail + if cause != nil { + detail = boundedError(cause) + } + payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(&rvboxv1.CommandEvent{IssueUuid: issue.String(), ObservedAt: timestamppb.New(executor.Now()), Payload: &rvboxv1.CommandEvent_SignalResult{SignalResult: &rvboxv1.SignalResult{Signal: signal, Accepted: accepted, GracefulDeliveryAttempted: signal == rvboxv1.SignalKind_SIGNAL_TERM, ForcedTerminationUsed: outcome.Escalated || signal == rvboxv1.SignalKind_SIGNAL_KILL, Detail: detail, CommandRevision: revision}}}) + if err != nil { + return err + } + _, err = executor.Store.AppendEvent(ctx, issue, spool.EventInput{Kind: 8, Compression: 1, RawBytes: uint64(len(payload)), Payload: payload, CreatedAt: executor.Now()}) + if err == nil { + executor.notify(issue) + } + return err +} + +func (executor *Executor) Terminate(ctx context.Context, issue domain.UUID) error { + executor.mu.Lock() + process := executor.active[issue] + executor.forced[issue] = true + executor.mu.Unlock() + if process == nil { + return nil + } + returnError := error(nil) + if _, err := executor.Supervisor.Signal(ctx, process, supervisor.SignalKill); err != nil { + returnError = err + } + return returnError +} + +func (executor *Executor) StopAll(ctx context.Context) error { + return executor.Supervisor.StopAll(ctx) +} + +func (executor *Executor) remove(issue domain.UUID) { + executor.mu.Lock() + delete(executor.active, issue) + delete(executor.cancel, issue) + delete(executor.forced, issue) + executor.mu.Unlock() +} + +func (executor *Executor) notify(issue domain.UUID) { + if executor.Notify != nil { + executor.Notify(issue) + } +} + +func cloneEnvironment(input map[string]string) map[string]string { + if input == nil { + return nil + } + result := make(map[string]string, len(input)) + for key, value := range input { + result[key] = value + } + return result +} + +func executionProfileNames(input []rvboxv1.ExecutionProfile) []string { + result := make([]string, 0, len(input)) + for _, profile := range input { + result = append(result, profile.String()) + } + return result +} diff --git a/internal/client/agent/executor_test.go b/internal/client/agent/executor_test.go new file mode 100644 index 0000000..0a7848a --- /dev/null +++ b/internal/client/agent/executor_test.go @@ -0,0 +1,135 @@ +//go:build !windows + +package agent + +import ( + "context" + "crypto/sha256" + "path/filepath" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/client/spool" + clientwindows "github.com/rvbox/rvbox/internal/client/supervisor/windows" + "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/timestamppb" +) + +func TestExecutorRunsDurableCommandAndPublishesTerminalEvents_HP_EXECUTOR_01(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + store, err := spool.Open(ctx, spool.Options{DataDir: filepath.Join(t.TempDir(), "spool"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer store.Close() + issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-0000000000d1") + spec := &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_SH, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "printf executor-ok"}} + encoded, err := proto.Marshal(spec) + if err != nil { + t.Fatal(err) + } + requestHash := sha256.Sum256([]byte("executor-request")) + if _, err := store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: requestHash, Revision: 1, Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED), ExecutionSpec: encoded}, time.Now().UTC()); err != nil { + t.Fatal(err) + } + notified := make(chan domain.UUID, 16) + supervised, err := clientwindows.NewSupervisor(clientwindows.NativeOptions{MaxOutputChunk: 64}) + if err != nil { + t.Fatal(err) + } + executor, err := NewExecutor(ExecutorOptions{Store: store, Supervisor: supervised, WorkDir: t.TempDir(), Notify: func(value domain.UUID) { notified <- value }}) + if err != nil { + t.Fatal(err) + } + dispatch := &rvboxv1.CommandDispatch{IssueUuid: issue.String(), CommandRevision: 1, TargetSessionGeneration: 1, IssueTime: timestamppb.Now(), Spec: spec, ImmutableRequestSha256: requestHash[:]} + if err := executor.Dispatch(ctx, Session{Generation: 1}, dispatch); err != nil { + t.Fatal(err) + } + deadline := time.After(5 * time.Second) + for { + command, err := store.GetCommand(ctx, issue) + if err != nil { + t.Fatal(err) + } + if command.Phase == uint32(rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED) { + break + } + select { + case <-deadline: + t.Fatalf("executor did not reach terminal phase: %d", command.Phase) + case <-notified: + } + } + if _, err := store.AssignSendWindow(ctx, issue, 8, 1<<20); err != nil { + t.Fatal(err) + } + events, err := store.PendingEvents(ctx, issue) + if err != nil || len(events) < 3 { + t.Fatalf("durable executor events = %#v, %v", events, err) + } + if _, err := store.Check(ctx); err != nil { + t.Fatal(err) + } +} + +func TestExecutorWaitsForScriptCommitBeforeLaunch_HP_EXECUTOR_02(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + store, err := spool.Open(ctx, spool.Options{DataDir: filepath.Join(t.TempDir(), "spool"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer store.Close() + issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-0000000000d2") + body := []byte("printf script-ok") + digest := sha256.Sum256(body) + spec := &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_SH, Source: &rvboxv1.ExecutionSpec_Script{Script: &rvboxv1.ScriptDescriptor{Filename: "ignored.sh", SizeBytes: uint64(len(body)), Sha256: digest[:]}}} + encoded, _ := proto.Marshal(spec) + hash := sha256.Sum256([]byte("script-executor-request")) + if _, err := store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: hash, Revision: 1, Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED), ExecutionSpec: encoded, Script: &spool.ScriptDescriptor{SizeBytes: uint64(len(body)), SHA256: digest}}, time.Now().UTC()); err != nil { + t.Fatal(err) + } + supervised, _ := clientwindows.NewSupervisor(clientwindows.NativeOptions{}) + executor, err := NewExecutor(ExecutorOptions{Store: store, Supervisor: supervised, WorkDir: t.TempDir()}) + if err != nil { + t.Fatal(err) + } + dispatch := &rvboxv1.CommandDispatch{IssueUuid: issue.String(), CommandRevision: 1, TargetSessionGeneration: 1, Spec: spec, ImmutableRequestSha256: hash[:]} + if err := executor.Dispatch(ctx, Session{Generation: 1}, dispatch); err != nil { + t.Fatal(err) + } + command, _ := store.GetCommand(ctx, issue) + if command.Phase != uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED) { + t.Fatalf("script launched before commit, phase=%d", command.Phase) + } + chunkHash := sha256.Sum256(body) + if _, err := store.AppendScriptChunk(ctx, issue, 0, body, chunkHash); err != nil { + t.Fatal(err) + } + if _, err := store.CommitScript(ctx, issue, spool.ScriptDescriptor{SizeBytes: uint64(len(body)), SHA256: digest}); err != nil { + t.Fatal(err) + } + if err := executor.ScriptReady(ctx, Session{Generation: 1}, issue); err != nil { + t.Fatal(err) + } + deadline := time.After(5 * time.Second) + for { + command, err = store.GetCommand(ctx, issue) + if err != nil { + t.Fatal(err) + } + if command.Phase == uint32(rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED) { + return + } + select { + case <-deadline: + t.Fatal("committed script did not reach terminal phase") + case <-time.After(10 * time.Millisecond): + } + } +} diff --git a/internal/client/agent/handshake.go b/internal/client/agent/handshake.go index 66c607a..2c4b6e3 100644 --- a/internal/client/agent/handshake.go +++ b/internal/client/agent/handshake.go @@ -77,16 +77,27 @@ func Reconcile(ctx context.Context, transport Transport, session Session, snapsh if err := transport.Write(ctx, encoded); err != nil { return nil, err } - received, err := transport.Read(ctx) - if err != nil { - return nil, err + for { + received, err := transport.Read(ctx) + if err != nil { + return nil, err + } + result, err := agentproto.DecodeEnvelope(received, limits, rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil { + return nil, fmt.Errorf("decode reconciliation response: %w", err) + } + if result.GetSessionId() != session.ID || result.GetSessionGeneration() != session.Generation { + return nil, ErrProtocolHandshake + } + if result.GetReconcileRequest() != nil { + // The request may have been queued before the client's snapshot write + // reached the server. The snapshot is already on the wire; consume the + // advisory target frame and continue waiting for the durable result. + continue + } + if result.GetReconcileResult() == nil { + return nil, ErrProtocolHandshake + } + return proto.Clone(result.GetReconcileResult()).(*rvboxv1.ReconcileResult), nil } - result, err := agentproto.DecodeEnvelope(received, limits, rvboxv1.Platform_PLATFORM_WINDOWS) - if err != nil { - return nil, fmt.Errorf("decode ReconcileResult: %w", err) - } - if result.GetSessionId() != session.ID || result.GetSessionGeneration() != session.Generation || result.GetReconcileResult() == nil { - return nil, ErrProtocolHandshake - } - return proto.Clone(result.GetReconcileResult()).(*rvboxv1.ReconcileResult), nil } diff --git a/internal/client/agent/runner.go b/internal/client/agent/runner.go new file mode 100644 index 0000000..f796a0f --- /dev/null +++ b/internal/client/agent/runner.go @@ -0,0 +1,543 @@ +package agent + +import ( + "context" + "crypto/rand" + "errors" + "fmt" + "math/big" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "github.com/rvbox/rvbox/internal/client/spool" + "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/timestamppb" +) + +// RunnerOptions contains the replaceable network edge and the durable client +// state used by one daemon. The supervisor is deliberately a callback here: +// process execution can outlive a network session and is owned by the caller. +type RunnerOptions struct { + Store *spool.Store + Dial func(context.Context) (Transport, error) + Hello *rvboxv1.ClientHello + Limits agentproto.Limits + + Backoff BackoffOptions + Jitter Jitter + Now func() time.Time + + OnDispatch func(context.Context, Session, *rvboxv1.CommandDispatch) error + OnStdin func(context.Context, Session, *rvboxv1.StdinWrite) error + OnCloseStdin func(context.Context, Session, *rvboxv1.CloseStdin) error + OnSignal func(context.Context, Session, *rvboxv1.SignalCommand) error + OnScriptReady func(context.Context, Session, domain.UUID) error + OnTerminate func(context.Context, domain.UUID) error + // EventReady wakes the active session after a supervisor worker appends a + // durable event. The network loop remains the sole writer; a reconnect can + // safely ignore a stale notification because replay reads the spool again. + EventReady <-chan domain.UUID +} + +// BackoffOptions and Jitter are kept at the agent boundary so callers do not +// need to depend on the internal session-machine implementation. +type BackoffOptions struct { + Initial time.Duration + Maximum time.Duration + StableReset time.Duration +} + +func (options BackoffOptions) Validate() error { + if options.Initial <= 0 || options.Maximum < options.Initial || options.StableReset <= 0 { + return errors.New("invalid client reconnect backoff options") + } + return nil +} + +type Jitter func(time.Duration) time.Duration + +func (options RunnerOptions) validate() error { + if options.Store == nil || options.Dial == nil || options.Hello == nil { + return errors.New("client runner requires store, dialer, and hello") + } + if options.Limits.MaxEnvelopeBytes == 0 { + options.Limits = agentproto.DefaultLimits() + } + if options.Backoff.Initial == 0 { + options.Backoff = BackoffOptions{Initial: time.Second, Maximum: time.Minute, StableReset: time.Minute} + } + if err := options.Backoff.Validate(); err != nil { + return err + } + if options.Jitter == nil { + return errors.New("client runner jitter is required") + } + if options.Now == nil { + return errors.New("client runner clock is required") + } + return nil +} + +// Run reconnects until ctx is cancelled. A failed handshake or transport is +// treated as a normal network failure; durable commands and their spool remain +// untouched. The first failure uses the configured initial backoff and a +// successful session resets the exponential history only after StableReset. +func Run(ctx context.Context, options RunnerOptions) error { + if err := options.validate(); err != nil { + return err + } + delay := time.Duration(0) + failures := uint32(0) + for { + if err := waitUntil(ctx, delay); err != nil { + return nil + } + sessionStarted := time.Now() + _ = runOnce(ctx, options) + if ctx.Err() != nil { + return nil + } + if time.Since(sessionStarted) >= options.Backoff.StableReset { + failures = 0 + } + if failures < ^uint32(0) { + failures++ + } + cap := options.Backoff.Initial + for index := uint32(1); index < failures && cap < options.Backoff.Maximum; index++ { + if cap > options.Backoff.Maximum/2 { + cap = options.Backoff.Maximum + break + } + cap *= 2 + } + if cap > options.Backoff.Maximum { + cap = options.Backoff.Maximum + } + delay = options.Jitter(cap) + if delay < 0 { + delay = 0 + } + if delay > cap { + delay = cap + } + } +} + +func waitUntil(ctx context.Context, delay time.Duration) error { + if delay <= 0 { + return nil + } + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} + +// RunOnce performs one complete session and returns when the transport is +// lost or a protocol/storage error makes that session unusable. +func RunOnce(ctx context.Context, options RunnerOptions) error { + if err := options.validate(); err != nil { + return err + } + return runOnce(ctx, options) +} + +func runOnce(ctx context.Context, options RunnerOptions) (resultErr error) { + transport, err := options.Dial(ctx) + if err != nil { + return err + } + defer func() { + if closeErr := transport.Close(); resultErr == nil && closeErr != nil { + resultErr = closeErr + } + }() + + limits := options.Limits + if limits.MaxEnvelopeBytes == 0 { + limits = agentproto.DefaultLimits() + } + session, err := Handshake(ctx, transport, options.Hello, limits) + if err != nil { + return err + } + snapshot, err := options.Store.ReconcileSnapshot(ctx) + if err != nil { + return fmt.Errorf("build client reconciliation snapshot: %w", err) + } + result, err := Reconcile(ctx, transport, session, snapshot, limits) + if err != nil { + return err + } + terminated, err := ApplyReconcileResult(ctx, options.Store, result, options.Now()) + if err != nil { + return fmt.Errorf("apply server reconciliation: %w", err) + } + for _, issue := range terminated { + if options.OnTerminate != nil { + if err := options.OnTerminate(ctx, issue); err != nil { + return err + } + } + } + sentEvents := make(map[domain.UUID]uint64) + if err := replayEvents(ctx, transport, options.Store, session, snapshot, limits, sentEvents); err != nil { + return err + } + if err := sendCapacity(ctx, transport, session, options.Hello, snapshot, limits); err != nil { + return err + } + return serveActive(ctx, transport, options, session, sentEvents) +} + +func replayEvents(ctx context.Context, transport Transport, store *spool.Store, session Session, snapshot *rvboxv1.ReconcileSnapshot, limits agentproto.Limits, sent map[domain.UUID]uint64) error { + for _, retained := range snapshot.GetRetainedCommands() { + if retained == nil || retained.GetTombstoned() { + continue + } + issue, err := domain.ParseUUIDv7(retained.GetIssueUuid()) + if err != nil { + return err + } + if _, err := store.AssignSendWindow(ctx, issue, 128, 1<<20); err != nil { + return err + } + events, err := store.PendingEvents(ctx, issue) + if err != nil { + return err + } + for _, event := range events { + if err := SendStoredEvent(ctx, transport, session, event, limits); err != nil { + return err + } + if event.EventSeq > sent[issue] { + sent[issue] = event.EventSeq + } + } + } + return nil +} + +// SendStoredEvent converts the compact local representation back to a +// canonical CommandEvent. Output rows store only their typed OutputChunk to +// avoid duplicating the envelope; metadata rows may store a full event. +func SendStoredEvent(ctx context.Context, transport Transport, session Session, event spool.Event, limits agentproto.Limits) error { + commandEvent, err := commandEventFromStored(event) + if err != nil { + return err + } + commandEvent.IssueUuid = event.IssueUUID.String() + commandEvent.EventSeq = event.EventSeq + commandEvent.ObservedAt = timestamppb.New(event.CreatedAt) + return SendCommandEvent(ctx, transport, session, commandEvent, limits) +} + +func commandEventFromStored(event spool.Event) (*rvboxv1.CommandEvent, error) { + if event.EventSeq == 0 || event.IssueUUID == (domain.UUID{}) || event.CreatedAt.IsZero() { + return nil, errors.New("stored event lacks assigned sequence or owner") + } + if event.Kind == spool.EventKindOutput { + var output rvboxv1.OutputChunk + if err := proto.Unmarshal(event.Payload, &output); err != nil { + return nil, fmt.Errorf("decode stored output event: %w", err) + } + return &rvboxv1.CommandEvent{Payload: &rvboxv1.CommandEvent_Output{Output: &output}}, nil + } + if event.Kind == spool.EventKindOutputTruncation { + var marker rvboxv1.OutputTruncation + if err := proto.Unmarshal(event.Payload, &marker); err != nil { + return nil, fmt.Errorf("decode stored truncation event: %w", err) + } + return &rvboxv1.CommandEvent{Payload: &rvboxv1.CommandEvent_OutputTruncation{OutputTruncation: &marker}}, nil + } + var stored rvboxv1.CommandEvent + if err := proto.Unmarshal(event.Payload, &stored); err != nil || stored.Payload == nil { + return nil, fmt.Errorf("decode stored command event: %w", err) + } + return &stored, nil +} + +func sendCapacity(ctx context.Context, transport Transport, session Session, hello *rvboxv1.ClientHello, snapshot *rvboxv1.ReconcileSnapshot, limits agentproto.Limits) error { + var running, queued uint32 + for _, retained := range snapshot.GetRetainedCommands() { + if retained == nil || retained.GetTombstoned() { + continue + } + switch retained.GetLifecycle() { + case rvboxv1.CommandLifecycle_COMMAND_ACCEPTED, rvboxv1.CommandLifecycle_COMMAND_RUNNING: + running++ + case rvboxv1.CommandLifecycle_COMMAND_QUEUED, rvboxv1.CommandLifecycle_COMMAND_DISPATCHED: + queued++ + } + } + if running > hello.GetMaxRunningCommands() || queued > hello.GetMaxQueuedCommands() { + return fmt.Errorf("durable client command count exceeds advertised capacity") + } + envelope := &rvboxv1.AgentEnvelope{SessionId: session.ID, SessionGeneration: session.Generation, Payload: &rvboxv1.AgentEnvelope_ClientCapacity{ClientCapacity: &rvboxv1.ClientCapacity{RunningCommands: running, QueuedCommands: queued, MaxRunningCommands: hello.GetMaxRunningCommands(), MaxQueuedCommands: hello.GetMaxQueuedCommands()}}} + if err := agentproto.ValidateEnvelope(envelope, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err != nil { + return err + } + encoded, err := proto.Marshal(envelope) + if err != nil { + return err + } + return transport.Write(ctx, encoded) +} + +func serveActive(ctx context.Context, transport Transport, options RunnerOptions, session Session, sent map[domain.UUID]uint64) error { + limits := options.Limits + if limits.MaxEnvelopeBytes == 0 { + limits = agentproto.DefaultLimits() + } + readContext, cancelRead := context.WithCancel(ctx) + defer cancelRead() + frames := make(chan []byte, 1) + readErrors := make(chan error, 1) + go func() { + for { + encoded, err := transport.Read(readContext) + if err != nil { + select { + case readErrors <- err: + case <-readContext.Done(): + } + return + } + select { + case frames <- encoded: + case <-readContext.Done(): + return + } + } + }() + for { + var encoded []byte + select { + case <-ctx.Done(): + return nil + case err := <-readErrors: + return err + case issue := <-options.EventReady: + if issue == (domain.UUID{}) { + continue + } + if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil { + return err + } + continue + case encoded = <-frames: + } + envelope, err := agentproto.DecodeEnvelope(encoded, limits, rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil { + return err + } + if envelope.GetSessionId() != session.ID || envelope.GetSessionGeneration() != session.Generation { + return ErrProtocolHandshake + } + switch { + case envelope.GetCommandDispatch() != nil: + if err := handleDispatch(ctx, transport, options, session, envelope.GetCommandDispatch(), limits); err != nil { + return err + } + issue, parseErr := domain.ParseUUIDv7(envelope.GetCommandDispatch().GetIssueUuid()) + if parseErr == nil { + if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil { + return err + } + } + case envelope.GetEventAck() != nil: + if err := ApplyEventAck(ctx, options.Store, envelope.GetEventAck()); err != nil { + return err + } + case envelope.GetScriptChunk() != nil: + if err := handleScriptChunk(ctx, transport, options.Store, session, envelope.GetScriptChunk(), limits); err != nil { + return err + } + case envelope.GetScriptCommit() != nil: + if err := handleScriptCommit(ctx, transport, options.Store, session, envelope.GetScriptCommit(), limits); err != nil { + return err + } + if options.OnScriptReady != nil { + issue, err := domain.ParseUUIDv7(envelope.GetScriptCommit().GetIssueUuid()) + if err != nil { + return err + } + if err := options.OnScriptReady(ctx, session, issue); err != nil { + return err + } + } + case envelope.GetStdinWrite() != nil: + if options.OnStdin != nil { + if err := options.OnStdin(ctx, session, envelope.GetStdinWrite()); err != nil { + return err + } + } + if issue, parseErr := domain.ParseUUIDv7(envelope.GetStdinWrite().GetIssueUuid()); parseErr == nil { + if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil { + return err + } + } + case envelope.GetCloseStdin() != nil: + if options.OnCloseStdin != nil { + if err := options.OnCloseStdin(ctx, session, envelope.GetCloseStdin()); err != nil { + return err + } + } + if issue, parseErr := domain.ParseUUIDv7(envelope.GetCloseStdin().GetIssueUuid()); parseErr == nil { + if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil { + return err + } + } + case envelope.GetSignalCommand() != nil: + if options.OnSignal != nil { + if err := options.OnSignal(ctx, session, envelope.GetSignalCommand()); err != nil { + return err + } + } + if issue, parseErr := domain.ParseUUIDv7(envelope.GetSignalCommand().GetIssueUuid()); parseErr == nil { + if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil { + return err + } + } + case envelope.GetError() != nil: + if envelope.GetError().GetCloseSession() { + return fmt.Errorf("server closed agent session: %s", envelope.GetError().GetError().GetMessage()) + } + case envelope.GetReconcileRequest() != nil: + // A request is advisory after the snapshot/result barrier. The next + // reconnect repeats the complete snapshot; never apply partial targets. + default: + return ErrUnexpectedMessage + } + } +} + +func flushEvents(ctx context.Context, transport Transport, store *spool.Store, session Session, issue domain.UUID, limits agentproto.Limits, sent map[domain.UUID]uint64) error { + if issue == (domain.UUID{}) { + return nil + } + if _, err := store.AssignSendWindow(ctx, issue, 128, 1<<20); err != nil { + return err + } + events, err := store.PendingEvents(ctx, issue) + if err != nil { + return err + } + for _, event := range events { + if event.EventSeq == 0 || event.EventSeq <= sent[issue] { + continue + } + if err := SendStoredEvent(ctx, transport, session, event, limits); err != nil { + return err + } + sent[issue] = event.EventSeq + } + return nil +} + +func handleDispatch(ctx context.Context, transport Transport, options RunnerOptions, session Session, dispatch *rvboxv1.CommandDispatch, limits agentproto.Limits) error { + acceptance, err := PersistDispatch(ctx, options.Store, session, dispatch, options.Now(), limits) + accepted := err == nil + ack := &rvboxv1.CommandAccepted{IssueUuid: dispatch.GetIssueUuid(), CommandRevision: dispatch.GetCommandRevision(), Accepted: accepted} + if err != nil { + ack.Rejection = rejectionForError(err, dispatch.GetIssueUuid()) + } + envelope := &rvboxv1.AgentEnvelope{SessionId: session.ID, SessionGeneration: session.Generation, Payload: &rvboxv1.AgentEnvelope_CommandAccepted{CommandAccepted: ack}} + if validateErr := agentproto.ValidateEnvelope(envelope, limits, rvboxv1.Platform_PLATFORM_WINDOWS); validateErr != nil { + return validateErr + } + encoded, marshalErr := proto.Marshal(envelope) + if marshalErr != nil { + return marshalErr + } + if err := transport.Write(ctx, encoded); err != nil { + return err + } + if accepted && options.OnDispatch != nil { + return options.OnDispatch(ctx, session, dispatch) + } + _ = acceptance + return nil +} + +func rejectionForError(err error, issue string) *rvboxv1.ControlError { + code := rvboxv1.ControlError_TRANSIENT + if errors.Is(err, spool.ErrCommandConflict) || errors.Is(err, spool.ErrAlreadyExecuted) { + code = rvboxv1.ControlError_CONFLICT + } else if errors.Is(err, agentproto.ErrInvalidExecutionSpec) || errors.Is(err, ErrProtocolHandshake) { + code = rvboxv1.ControlError_INVALID_ARGUMENT + } + return &rvboxv1.ControlError{Code: code, Message: boundedError(err), Retryable: code == rvboxv1.ControlError_TRANSIENT, IssueUuid: issue} +} + +func boundedError(err error) string { + if err == nil { + return "" + } + message := err.Error() + if len(message) > 1024 { + message = message[:1024] + } + return message +} + +func handleScriptChunk(ctx context.Context, transport Transport, store *spool.Store, session Session, chunk *rvboxv1.ScriptChunk, limits agentproto.Limits) error { + status, err := ApplyScriptChunk(ctx, store, session, chunk, limits) + if err != nil { + return err + } + return appendAndSendScriptStatus(ctx, transport, store, session, chunk.GetIssueUuid(), status, limits) +} + +func handleScriptCommit(ctx context.Context, transport Transport, store *spool.Store, session Session, commit *rvboxv1.ScriptCommit, limits agentproto.Limits) error { + status, err := ApplyScriptCommit(ctx, store, session, commit, limits) + if err != nil { + return err + } + return appendAndSendScriptStatus(ctx, transport, store, session, commit.GetIssueUuid(), status, limits) +} + +func appendAndSendScriptStatus(ctx context.Context, transport Transport, store *spool.Store, session Session, issueText string, status spool.ScriptStatus, limits agentproto.Limits) error { + issue, err := domain.ParseUUIDv7(issueText) + if err != nil { + return err + } + event := &rvboxv1.CommandEvent{IssueUuid: issueText, ObservedAt: timestamppb.New(time.Now().UTC()), Payload: &rvboxv1.CommandEvent_ScriptStatus{ScriptStatus: &rvboxv1.ScriptUploadStatus{ReceivedBytes: status.ReceivedBytes, Complete: status.Committed}}} + payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(event) + if err != nil { + return err + } + if _, err := store.AppendEvent(ctx, issue, spool.EventInput{Kind: 9, Compression: 1, RawBytes: uint64(len(payload)), Payload: payload, CreatedAt: event.GetObservedAt().AsTime()}); err != nil { + return err + } + assigned, err := store.AssignSendWindow(ctx, issue, 1, 1<<20) + if err != nil { + return err + } + for _, item := range assigned { + if err := SendStoredEvent(ctx, transport, session, item, limits); err != nil { + return err + } + } + return nil +} + +// CryptoJitter returns a full-jitter delay without relying on math/rand's +// process-global state. Tests inject a deterministic jitter instead. +func CryptoJitter(capacity time.Duration) time.Duration { + if capacity <= 0 { + return 0 + } + value, err := rand.Int(rand.Reader, big.NewInt(int64(capacity)+1)) + if err != nil { + return capacity + } + return time.Duration(value.Int64()) +} diff --git a/internal/client/agent/runner_test.go b/internal/client/agent/runner_test.go new file mode 100644 index 0000000..60f9ca7 --- /dev/null +++ b/internal/client/agent/runner_test.go @@ -0,0 +1,139 @@ +package agent + +import ( + "context" + "errors" + "path/filepath" + "sync" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "github.com/rvbox/rvbox/internal/client/spool" + "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/timestamppb" +) + +func TestRunOnceReconcilesReplaysAndAdvertisesCapacity_HP_RUNTIME_01(t *testing.T) { + ctx := context.Background() + store, err := spool.Open(ctx, spool.Options{DataDir: filepath.Join(t.TempDir(), "spool"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + issue, err := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-0000000000b1") + if err != nil { + t.Fatal(err) + } + hash := [32]byte{1} + if _, err := store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: hash, Revision: 1, Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED), ExecutionSpec: []byte("spec")}, time.Now().UTC()); err != nil { + t.Fatal(err) + } + eventPayload, err := proto.Marshal(&rvboxv1.CommandEvent{Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: &rvboxv1.LifecycleChange{Lifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, CommandRevision: 1}}}) + if err != nil { + t.Fatal(err) + } + if _, err := store.AppendEvent(ctx, issue, spool.EventInput{Kind: 4, Compression: 1, RawBytes: uint64(len(eventPayload)), Payload: eventPayload, CreatedAt: time.Now().UTC()}); err != nil { + t.Fatal(err) + } + if _, err := store.AssignSendWindow(ctx, issue, 1, 1<<20); err != nil { + t.Fatal(err) + } + welcome, _ := proto.Marshal(&rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ServerWelcome{ServerWelcome: &rvboxv1.ServerWelcome{SelectedProtocol: &rvboxv1.ProtocolVersion{Major: 1, Minor: 0}, ServerTime: timestamppb.Now()}}}) + reconcileRequest, _ := proto.Marshal(&rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ReconcileRequest{ReconcileRequest: &rvboxv1.ReconcileRequest{}}}) + reconcileResult, _ := proto.Marshal(&rvboxv1.AgentEnvelope{SessionId: "session", SessionGeneration: 1, Payload: &rvboxv1.AgentEnvelope_ReconcileResult{ReconcileResult: &rvboxv1.ReconcileResult{}}}) + transport := &runnerTransport{reads: [][]byte{welcome, reconcileRequest, reconcileResult}, terminal: errors.New("transport closed")} + hello := &rvboxv1.ClientHello{ClientId: "runner-client", SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, DaemonVersion: "test", Platform: rvboxv1.Platform_PLATFORM_WINDOWS, Architecture: "amd64", DaemonCwd: `C:\`, SupportedShells: []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_POWERSHELL}, ClientInstanceId: store.ClientInstanceID().String(), MaxRunningCommands: 1, MaxQueuedCommands: 1, SentAt: timestamppb.Now()} + err = RunOnce(ctx, RunnerOptions{Store: store, Dial: func(context.Context) (Transport, error) { return transport, nil }, Hello: hello, Limits: agentproto.DefaultLimits(), Backoff: BackoffOptions{Initial: time.Millisecond, Maximum: time.Millisecond, StableReset: time.Second}, Jitter: func(value time.Duration) time.Duration { return 0 }, Now: func() time.Time { return time.Now().UTC() }}) + if !errors.Is(err, transport.terminal) { + t.Fatalf("RunOnce error = %v, want transport close", err) + } + transport.mu.Lock() + writes := append([][]byte(nil), transport.writes...) + transport.mu.Unlock() + if len(writes) != 4 { + t.Fatalf("wire write count = %d, want hello/snapshot/event/capacity", len(writes)) + } + last, err := agentproto.DecodeEnvelope(writes[len(writes)-1], agentproto.DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil || last.GetClientCapacity() == nil || last.GetClientCapacity().GetMaxRunningCommands() != 1 { + t.Fatalf("capacity advertisement = %#v, %v", last, err) + } + eventEnvelope, err := agentproto.DecodeEnvelope(writes[2], agentproto.DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil || eventEnvelope.GetCommandEvent() == nil || eventEnvelope.GetCommandEvent().GetEventSeq() != 1 || eventEnvelope.GetCommandEvent().GetIssueUuid() != issue.String() { + t.Fatalf("replayed event = %#v, %v", eventEnvelope, err) + } +} + +func TestRunnerOptionsRejectMissingJitter_HP_RUNTIME_02(t *testing.T) { + if err := (RunnerOptions{}).validate(); err == nil { + t.Fatal("empty runner options unexpectedly validated") + } +} + +func TestFlushEventsAssignsAndSendsOnlyUnacknowledgedRows_HP_RUNTIME_03(t *testing.T) { + ctx := context.Background() + store, err := spool.Open(ctx, spool.Options{DataDir: filepath.Join(t.TempDir(), "spool"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer store.Close() + issue, err := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-0000000000b2") + if err != nil { + t.Fatal(err) + } + hash := [32]byte{2} + if _, err := store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: hash, Revision: 1, Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED)}, time.Now().UTC()); err != nil { + t.Fatal(err) + } + payload, err := proto.Marshal(&rvboxv1.CommandEvent{Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: &rvboxv1.LifecycleChange{Lifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, CommandRevision: 1}}}) + if err != nil { + t.Fatal(err) + } + if _, err := store.AppendEvent(ctx, issue, spool.EventInput{Kind: 4, Compression: 1, RawBytes: uint64(len(payload)), Payload: payload, CreatedAt: time.Now().UTC()}); err != nil { + t.Fatal(err) + } + transport := &runnerTransport{} + sent := map[domain.UUID]uint64{} + session := Session{ID: "session", Generation: 1} + if err := flushEvents(ctx, transport, store, session, issue, agentproto.DefaultLimits(), sent); err != nil { + t.Fatal(err) + } + if sent[issue] != 1 || len(transport.writes) != 1 { + t.Fatalf("sent cursor/writes = %d/%d", sent[issue], len(transport.writes)) + } + if err := flushEvents(ctx, transport, store, session, issue, agentproto.DefaultLimits(), sent); err != nil { + t.Fatal(err) + } + if len(transport.writes) != 1 { + t.Fatalf("already-sent event was duplicated: %d writes", len(transport.writes)) + } +} + +type runnerTransport struct { + mu sync.Mutex + reads [][]byte + writes [][]byte + terminal error +} + +func (transport *runnerTransport) Write(_ context.Context, payload []byte) error { + transport.mu.Lock() + defer transport.mu.Unlock() + transport.writes = append(transport.writes, append([]byte(nil), payload...)) + return nil +} + +func (transport *runnerTransport) Read(context.Context) ([]byte, error) { + transport.mu.Lock() + defer transport.mu.Unlock() + if len(transport.reads) == 0 { + return nil, transport.terminal + } + payload := transport.reads[0] + transport.reads = transport.reads[1:] + return append([]byte(nil), payload...), nil +} + +func (transport *runnerTransport) Close() error { return nil } diff --git a/internal/client/spool/events.go b/internal/client/spool/events.go index 7640b3a..d2ffa01 100644 --- a/internal/client/spool/events.go +++ b/internal/client/spool/events.go @@ -8,9 +8,13 @@ import ( "errors" "fmt" "time" + "unicode/utf8" "github.com/klauspost/compress/zstd" + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/timestamppb" ) type Command struct { @@ -23,6 +27,11 @@ type Command struct { // server. It is retained as raw protobuf bytes so the runtime can validate // and execute exactly the admitted request after a restart. ExecutionSpec []byte + // Script carries the immutable descriptor reservation for a script-backed + // command. Its bytes arrive later through AppendScriptChunk; creating the + // descriptor row in the same transaction as acceptance prevents an + // accepted command from being left without upload state after a crash. + Script *ScriptDescriptor } type Acceptance struct { @@ -63,7 +72,7 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted return Acceptance{}, errors.New("invalid command acceptance") } var storedSpec []byte - var specCharge uint64 + var specCharge, scriptCharge uint64 if len(command.ExecutionSpec) > 0 { if uint64(len(command.ExecutionSpec)) > store.maxExecutionSpecBytes { return Acceptance{}, errors.New("execution specification exceeds client limit") @@ -78,6 +87,21 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted return Acceptance{}, err } } + var storedScript []byte + if command.Script != nil { + if err := validateAcceptedScript(command.ExecutionSpec, *command.Script, store.maxScriptBytes); err != nil { + return Acceptance{}, err + } + var err error + storedScript, err = compressScript(nil) + if err != nil { + return Acceptance{}, err + } + scriptCharge, err = EstimateCharge(ChargeInput{EncodedBytes: uint64(len(storedScript)), SQLiteRows: 1, IndexEntries: 1}) + if err != nil { + return Acceptance{}, err + } + } tx, err := store.db.BeginTx(ctx, nil) if err != nil { return Acceptance{}, err @@ -116,6 +140,9 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted return Acceptance{}, err } charge, overflow := addChecked(baseCharge, specCharge) + if !overflow { + charge, overflow = addChecked(charge, scriptCharge) + } if overflow { return Acceptance{}, &CapacityError{Tier: CapacityTierHardMaximum, Requested: ^uint64(0), Available: store.quotaLimits.HardAllocationBytes} } @@ -136,12 +163,32 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted return Acceptance{}, err } } + if command.Script != nil { + if _, err := tx.ExecContext(ctx, `INSERT INTO scripts(issue_uuid, declared_raw_bytes, declared_sha256, stored_bytes, compression, stored_data, charged_bytes) VALUES (?, ?, ?, ?, 2, ?, ?)`, command.IssueUUID[:], command.Script.SizeBytes, command.Script.SHA256[:], len(storedScript), storedScript, scriptCharge); err != nil { + return Acceptance{}, err + } + } if err := tx.Commit(); err != nil { return Acceptance{}, err } return Acceptance{Command: command}, nil } +func validateAcceptedScript(encodedSpec []byte, descriptor ScriptDescriptor, maximum uint64) error { + if descriptor.SizeBytes > maximum { + return ErrScriptBounds + } + var spec rvboxv1.ExecutionSpec + if err := proto.Unmarshal(encodedSpec, &spec); err != nil { + return ErrScriptConflict + } + declared := spec.GetScript() + if declared == nil || declared.GetSizeBytes() != descriptor.SizeBytes || len(declared.GetSha256()) != sha256.Size || !bytes.Equal(declared.GetSha256(), descriptor.SHA256[:]) { + return ErrScriptConflict + } + return nil +} + // GetCommand returns the durable command metadata and its immutable execution // specification. The returned protobuf bytes are a copy and can be decoded or // modified by the runtime without changing the spool's source of truth. @@ -340,6 +387,156 @@ func (store *Store) MarkTerminal(ctx context.Context, issueUUID domain.UUID, pha return nil } +// LaunchEvidence is the durable pre/post-authorization record used to fence +// uncertain OS launches across daemon restarts. PID is evidence only and is +// never sufficient for recovery-time signalling without a native creation +// identity check. +type LaunchEvidence struct { + Phase domain.LaunchPhase + Context string + PID uint32 +} + +func (store *Store) SetLaunchPhase(ctx context.Context, issueUUID domain.UUID, phase domain.LaunchPhase, contextName string, pid uint32) error { + if !validUUID(issueUUID) || phase > domain.LaunchPhaseAuthorized || len(contextName) > 128 || !utf8.ValidString(contextName) { + return errors.New("invalid launch barrier") + } + if phase == domain.LaunchPhaseNone { + contextName, pid = "", 0 + } + result, err := store.db.ExecContext(ctx, `UPDATE commands SET launch_phase = ?, launch_context = ?, launch_pid = ? WHERE issue_uuid = ? AND terminal = 0`, uint32(phase), contextName, pid, issueUUID[:]) + if err != nil { + return err + } + count, err := result.RowsAffected() + if err != nil { + return err + } + if count == 0 { + return ErrUnknownCommand + } + return nil +} + +// RecoverLaunchUncertainty converts every command whose authorization barrier +// crossed before a restart into one durable interrupted terminal event. The +// Windows Job kill-on-close guarantee makes this safe: a surviving process is +// never redispatched, and recovery never signals a PID by itself. +func (store *Store) RecoverLaunchUncertainty(ctx context.Context, now time.Time) ([]domain.UUID, error) { + if now.IsZero() { + return nil, errors.New("launch recovery time is required") + } + rows, err := store.db.QueryContext(ctx, `SELECT issue_uuid, command_revision FROM commands WHERE launch_phase = 2 AND terminal = 0 ORDER BY accepted_at, issue_uuid`) + if err != nil { + return nil, err + } + type pending struct { + issue domain.UUID + revision uint64 + } + var pendingRows []pending + for rows.Next() { + var encoded []byte + var revision uint64 + if err := rows.Scan(&encoded, &revision); err != nil { + _ = rows.Close() + return nil, err + } + var issue domain.UUID + if len(encoded) != len(issue) { + _ = rows.Close() + return nil, ErrScriptState + } + copy(issue[:], encoded) + if !validUUID(issue) || revision == 0 { + _ = rows.Close() + return nil, ErrScriptState + } + pendingRows = append(pendingRows, pending{issue: issue, revision: revision}) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + interrupted := make([]domain.UUID, 0, len(pendingRows)) + for _, item := range pendingRows { + if _, err := store.AppendLifecycle(ctx, item.issue, uint32(rvboxv1.CommandLifecycle_COMMAND_INTERRUPTED), item.revision, "uncertain launch recovered after daemon restart", now); err != nil { + return interrupted, err + } + interrupted = append(interrupted, item.issue) + } + return interrupted, nil +} + +// AppendLifecycle atomically advances the durable command phase and appends +// its lifecycle event. Keeping these writes in one transaction prevents a +// crash between a terminal marker and its public event from creating a state +// that can be replayed as a second execution. +func (store *Store) AppendLifecycle(ctx context.Context, issueUUID domain.UUID, phase uint32, revision uint64, detail string, observedAt time.Time) (Event, error) { + if !validUUID(issueUUID) || phase == 0 || phase > 11 || observedAt.IsZero() || revision == 0 { + return Event{}, errors.New("invalid lifecycle event") + } + if len(detail) > 4096 || !utf8.ValidString(detail) { + return Event{}, errors.New("lifecycle detail is invalid or too large") + } + tx, err := store.db.BeginTx(ctx, nil) + if err != nil { + return Event{}, err + } + defer tx.Rollback() + var current uint32 + var storedRevision uint64 + var nextOrdinal, outputCharged, totalCharged, closeout uint64 + if err := tx.QueryRowContext(ctx, `SELECT phase, command_revision, next_local_ordinal, output_charged_bytes, total_charged_bytes, closeout_remaining_bytes FROM commands WHERE issue_uuid = ?`, issueUUID[:]).Scan(¤t, &storedRevision, &nextOrdinal, &outputCharged, &totalCharged, &closeout); err == sql.ErrNoRows { + return Event{}, ErrUnknownCommand + } else if err != nil { + return Event{}, err + } + if storedRevision != revision { + return Event{}, ErrCommandConflict + } + if current == phase { + return Event{}, nil + } + if !domain.CanTransition(rvboxv1.CommandLifecycle(current), rvboxv1.CommandLifecycle(phase)) { + return Event{}, fmt.Errorf("invalid lifecycle transition %s -> %s", rvboxv1.CommandLifecycle(current), rvboxv1.CommandLifecycle(phase)) + } + lifecycle := &rvboxv1.LifecycleChange{Lifecycle: rvboxv1.CommandLifecycle(phase), CommandRevision: revision, Detail: detail} + payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(&rvboxv1.CommandEvent{IssueUuid: issueUUID.String(), ObservedAt: timestamppb.New(observedAt), Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: lifecycle}}) + if err != nil { + return Event{}, err + } + charge, err := EstimateCharge(ChargeInput{EncodedBytes: uint64(len(payload)), SQLiteRows: 1, IndexEntries: 2}) + if err != nil { + return Event{}, err + } + clientTotal, err := clientTotalCharge(ctx, tx) + if err != nil { + return Event{}, err + } + decision, err := CheckReservation(store.quotaLimits, ReservationState{CommandOutputCharged: outputCharged, CommandTotalCharged: totalCharged, ClientTotalCharged: clientTotal, CloseoutRemaining: closeout}, ReservationRequest{ChargedBytes: charge, UseCloseout: isTerminalPhase(phase)}) + if err != nil { + return Event{}, err + } + digest := immutableDigest(payload) + if _, err := tx.ExecContext(ctx, `INSERT INTO events(issue_uuid, local_ordinal, event_kind, compression, raw_bytes, charged_bytes, output, payload, payload_sha256, created_at) VALUES (?, ?, 4, 1, ?, ?, 0, ?, ?, ?)`, issueUUID[:], nextOrdinal, len(payload), charge, payload, digest[:], observedAt.UnixNano()); err != nil { + return Event{}, err + } + terminal := isTerminalPhase(phase) + if _, err := tx.ExecContext(ctx, `UPDATE commands SET phase = ?, terminal = ?, launch_phase = CASE WHEN ? = 1 THEN 0 ELSE launch_phase END, launch_context = CASE WHEN ? = 1 THEN '' ELSE launch_context END, launch_pid = CASE WHEN ? = 1 THEN 0 ELSE launch_pid END, next_local_ordinal = ?, total_charged_bytes = ?, closeout_remaining_bytes = ? WHERE issue_uuid = ?`, phase, boolInt(terminal), boolInt(terminal), boolInt(terminal), boolInt(terminal), nextOrdinal+1, decision.CommandTotalCharged, decision.CloseoutRemaining, issueUUID[:]); err != nil { + return Event{}, err + } + if err := updateClientTotalCharge(ctx, tx, decision.ClientTotalCharged); err != nil { + return Event{}, err + } + if err := tx.Commit(); err != nil { + return Event{}, err + } + return Event{IssueUUID: issueUUID, LocalOrdinal: nextOrdinal, Kind: 4, Compression: 1, RawBytes: uint64(len(payload)), Payload: append([]byte(nil), payload...), CreatedAt: observedAt}, nil +} + // CleanupTerminal moves a fully acknowledged terminal command into the compact // tombstone ledger and removes all command-owned spool data in one transaction. func (store *Store) CleanupTerminal(ctx context.Context, issueUUID domain.UUID, acknowledgedAt time.Time) error { diff --git a/internal/client/spool/migrations.go b/internal/client/spool/migrations.go index 8647322..58edae8 100644 --- a/internal/client/spool/migrations.go +++ b/internal/client/spool/migrations.go @@ -16,6 +16,7 @@ type migration struct { var migrations = []migration{ {version: 1, sql: schemaV1}, {version: 2, sql: schemaV2}, + {version: 3, sql: schemaV3}, } func applyMigrations(ctx context.Context, db *sql.DB) error { @@ -132,3 +133,15 @@ CREATE TABLE command_specs ( charged_bytes INTEGER NOT NULL CHECK(charged_bytes > 0) ) STRICT, WITHOUT ROWID; ` + +// schemaV3 adds a small durable launch barrier to every accepted command. A +// value of 2 means launch authorization may have crossed the OS boundary; on +// restart the client must interrupt that command instead of redispatching it. +// Keeping these fields on commands makes the barrier part of the existing +// command-owned quota/accounting row and lets terminal cleanup remove it with +// the command. +const schemaV3 = ` +ALTER TABLE commands ADD COLUMN launch_phase INTEGER NOT NULL DEFAULT 0 CHECK(launch_phase BETWEEN 0 AND 2); +ALTER TABLE commands ADD COLUMN launch_context TEXT NOT NULL DEFAULT ''; +ALTER TABLE commands ADD COLUMN launch_pid INTEGER NOT NULL DEFAULT 0 CHECK(launch_pid >= 0); +` diff --git a/internal/client/spool/recovery_test.go b/internal/client/spool/recovery_test.go index 2c66e9f..bbf4fd2 100644 --- a/internal/client/spool/recovery_test.go +++ b/internal/client/spool/recovery_test.go @@ -6,6 +6,9 @@ import ( "path/filepath" "testing" "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/domain" ) func TestCheckSpoolPayloadsAndCounters_HP_CLIENT_09(t *testing.T) { @@ -46,3 +49,33 @@ func TestCheckSpoolRejectsCounterDrift_BH_CLIENT_03(t *testing.T) { t.Fatalf("counter-drift Check error = %v, want ErrQuotaCounterMismatch", err) } } + +func TestRecoverLaunchUncertaintyFencesRedispatchAfterRestart_HP_CLIENT_12(t *testing.T) { + t.Parallel() + ctx := context.Background() + store := openTestStore(t, ctx, filepath.Join(t.TempDir(), "spool"), DefaultTombstoneLimit) + issue := testUUID(t, "019c46f1-1d02-7000-8000-000000000043") + command := testCommand(issue, []byte("uncertain launch")) + command.Phase = uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED) + now := time.Date(2026, time.September, 6, 12, 0, 0, 0, time.UTC) + if _, err := store.AcceptCommand(ctx, command, now); err != nil { + t.Fatal(err) + } + if err := store.SetLaunchPhase(ctx, issue, domain.LaunchPhaseAuthorized, "LOCAL_SYSTEM", 42); err != nil { + t.Fatal(err) + } + recovered, err := store.RecoverLaunchUncertainty(ctx, now.Add(time.Second)) + if err != nil || len(recovered) != 1 || recovered[0] != issue { + t.Fatalf("recovered = %v, %v", recovered, err) + } + var phase, launchPhase uint32 + if err := store.db.QueryRow(`SELECT phase, launch_phase FROM commands WHERE issue_uuid = ?`, issue[:]).Scan(&phase, &launchPhase); err != nil { + t.Fatal(err) + } + if phase != uint32(rvboxv1.CommandLifecycle_COMMAND_INTERRUPTED) || launchPhase != 0 { + t.Fatalf("recovered phase/barrier = %d/%d", phase, launchPhase) + } + if second, err := store.RecoverLaunchUncertainty(ctx, now.Add(2*time.Second)); err != nil || len(second) != 0 { + t.Fatalf("recovery repeated = %v, %v", second, err) + } +} diff --git a/internal/client/spool/script.go b/internal/client/spool/script.go index 8769318..bd748bd 100644 --- a/internal/client/spool/script.go +++ b/internal/client/spool/script.go @@ -20,6 +20,7 @@ var ( ErrScriptBounds = errors.New("script chunk is outside declared bounds") ErrScriptState = errors.New("stored script state is corrupt") ErrScriptTerminal = errors.New("terminal command cannot accept script data") + ErrScriptNotReady = errors.New("script has not been durably committed") ) type ScriptDescriptor struct { @@ -33,6 +34,39 @@ type ScriptStatus struct { Duplicate bool } +// ScriptBody returns a copy of the verified script body only after the +// contiguous upload has been committed. It is the sole spool read used by the +// supervisor; callers never reconstruct script bytes from individual chunks. +func (store *Store) ScriptBody(ctx context.Context, issueUUID domain.UUID) ([]byte, error) { + if !validUUID(issueUUID) { + return nil, ErrUnknownCommand + } + var row storedScript + var digest []byte + var storedBytes uint64 + var compression uint32 + var committed int + err := store.db.QueryRowContext(ctx, `SELECT declared_raw_bytes, declared_sha256, received_raw_bytes, stored_bytes, compression, stored_data, charged_bytes, committed FROM scripts WHERE issue_uuid = ?`, issueUUID[:]).Scan(&row.DeclaredBytes, &digest, &row.ReceivedBytes, &storedBytes, &compression, &row.Stored, &row.ChargedBytes, &committed) + if errors.Is(err, sql.ErrNoRows) { + return nil, ErrScriptNotReady + } + if err != nil { + return nil, err + } + if len(digest) != sha256.Size || storedBytes != uint64(len(row.Stored)) || compression != 2 || committed == 0 || row.ReceivedBytes != row.DeclaredBytes { + if committed == 0 { + return nil, ErrScriptNotReady + } + return nil, ErrScriptState + } + copy(row.DeclaredSHA256[:], digest) + body, err := decompressScript(row.Stored, row.ReceivedBytes, store.maxScriptBytes) + if err != nil || sha256.Sum256(body) != row.DeclaredSHA256 { + return nil, ErrScriptState + } + return append([]byte(nil), body...), nil +} + // BeginScript persists the immutable descriptor at command acceptance time. A // matching replay is harmless; a different descriptor is a protocol conflict. func (store *Store) BeginScript(ctx context.Context, issueUUID domain.UUID, descriptor ScriptDescriptor) (ScriptStatus, error) { diff --git a/internal/client/spool/spool_test.go b/internal/client/spool/spool_test.go index c8e176f..2fc15e1 100644 --- a/internal/client/spool/spool_test.go +++ b/internal/client/spool/spool_test.go @@ -10,6 +10,8 @@ import ( "testing" "time" + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/domain" ) @@ -186,6 +188,43 @@ func TestSpoolAcceptanceSequencingAndAck_HP_CLIENT_07(t *testing.T) { } } +func TestAppendLifecycleAtomicallyUpdatesPhaseAndEvent_HP_CLIENT_11(t *testing.T) { + t.Parallel() + ctx := context.Background() + store := openTestStore(t, ctx, filepath.Join(t.TempDir(), "spool"), DefaultTombstoneLimit) + issue := testUUID(t, "019c46f1-1d02-7000-8000-000000000021") + command := testCommand(issue, []byte("lifecycle")) + command.Phase = uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED) + now := time.Date(2026, time.September, 6, 12, 0, 0, 0, time.UTC) + if _, err := store.AcceptCommand(ctx, command, now); err != nil { + t.Fatal(err) + } + if _, err := store.AppendLifecycle(ctx, issue, uint32(rvboxv1.CommandLifecycle_COMMAND_RUNNING), 1, "launch authorized", now); err != nil { + t.Fatal(err) + } + if _, err := store.AppendLifecycle(ctx, issue, uint32(rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED), 1, "exit 0", now.Add(time.Second)); err != nil { + t.Fatal(err) + } + var phase uint32 + var terminal int + if err := store.db.QueryRow(`SELECT phase, terminal FROM commands WHERE issue_uuid = ?`, issue[:]).Scan(&phase, &terminal); err != nil { + t.Fatal(err) + } + if phase != uint32(rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED) || terminal != 1 { + t.Fatalf("phase/terminal = %d/%d", phase, terminal) + } + if _, err := store.AppendLifecycle(ctx, issue, uint32(rvboxv1.CommandLifecycle_COMMAND_RUNNING), 1, "illegal", now.Add(2*time.Second)); err == nil { + t.Fatal("terminal lifecycle regressed") + } + if _, err := store.AssignSendWindow(ctx, issue, 8, 1<<20); err != nil { + t.Fatal(err) + } + events, err := store.PendingEvents(ctx, issue) + if err != nil || len(events) != 2 { + t.Fatalf("lifecycle events = %#v, %v", events, err) + } +} + func TestTerminalCleanupTombstonesAndConflicts_BH_CLIENT_02(t *testing.T) { t.Parallel() ctx := context.Background() diff --git a/internal/client/supervisor/supervisor.go b/internal/client/supervisor/supervisor.go index 46d7811..da0faf7 100644 --- a/internal/client/supervisor/supervisor.go +++ b/internal/client/supervisor/supervisor.go @@ -28,9 +28,13 @@ const ( // still checks the identity/revision/source invariants because it is a second // durability boundary and may be called after a restart. type StartSpec struct { - IssueUUID domain.UUID - CommandRevision uint64 - Execution *rvboxv1.ExecutionSpec + IssueUUID domain.UUID + CommandRevision uint64 + Execution *rvboxv1.ExecutionSpec + // ScriptBody is the verified, durable body for Execution.script. It is + // supplied by the client spool only after the declared digest/length have + // been checked; command_text requests leave it empty. + ScriptBody []byte WorkingDirectory string Environment map[string]string ExecutionProfiles []string @@ -60,11 +64,23 @@ type EffectiveIdentity struct { type Process interface { IssueUUID() domain.UUID Identity() EffectiveIdentity + // ReadOutput returns the next bounded stdout/stderr chunk. It continues + // until both child pipes reach EOF, so Wait never reports a terminal + // result before the captured output has drained. + ReadOutput(context.Context) (OutputChunk, error) Wait(context.Context) (ExitStatus, error) WriteStdin(context.Context, []byte, bool) error CloseStdin(context.Context) error } +// OutputChunk is intentionally raw. Compression, quota admission, local +// ordering, and wire sequencing belong to the client spool rather than the +// operating-system supervisor. +type OutputChunk struct { + Stream rvboxv1.StreamKind + Data []byte +} + type ExitStatus struct { Code int32 Signaled bool diff --git a/internal/client/supervisor/windows/launch.go b/internal/client/supervisor/windows/launch.go index fa70357..a380d85 100644 --- a/internal/client/supervisor/windows/launch.go +++ b/internal/client/supervisor/windows/launch.go @@ -144,6 +144,8 @@ func BuildWrapper(spec *rvboxv1.ExecutionSpec, scriptBody []byte, maxBytes uint6 func shellTemplate(shellType rvboxv1.ShellType) (ShellPlan, error) { switch shellType { + case rvboxv1.ShellType_SHELL_SH, rvboxv1.ShellType_SHELL_BASH: + return ShellPlan{Type: shellType, WrapperExtension: ".sh", Encoding: WrapperEncodingUTF8}, nil case rvboxv1.ShellType_SHELL_CMD: return ShellPlan{Type: shellType, Arguments: []string{"/D", "/S", "/C"}, WrapperExtension: ".cmd", Encoding: WrapperEncodingUTF8}, nil case rvboxv1.ShellType_SHELL_POWERSHELL: @@ -192,6 +194,13 @@ func ValidateWindowsExecutablePath(path string) error { } } +// ValidAbsoluteWindowsPath reports whether value passes the lexical absolute +// path checks. It does not touch the filesystem; native callers must still +// re-stat the object and verify its ACL immediately before use. +func ValidAbsoluteWindowsPath(value string) bool { + return validAbsoluteWindowsPath(value) +} + func validAbsoluteWindowsPath(value string) bool { if value == "" || strings.IndexByte(value, 0) >= 0 || !utf8.ValidString(value) || strings.ContainsAny(value, "\r\n\t") { return false diff --git a/internal/client/supervisor/windows/native_exec.go b/internal/client/supervisor/windows/native_exec.go new file mode 100644 index 0000000..0d6a389 --- /dev/null +++ b/internal/client/supervisor/windows/native_exec.go @@ -0,0 +1,457 @@ +package windows + +// This file contains the platform-neutral process bookkeeping shared by the +// native Windows adapter and the deterministic non-Windows test adapter. The +// Windows build supplies the token/Job-backed start and signal operations in +// native_windows.go; keeping output and stdin semantics here prevents those +// paths from drifting apart. + +import ( + "context" + "errors" + "fmt" + "io" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/client/supervisor" + "github.com/rvbox/rvbox/internal/domain" +) + +var ( + ErrProcessAlreadyRunning = errors.New("a process for this command is already running") + ErrProcessNotFound = errors.New("supervised process was not found") + ErrProcessNotReady = errors.New("supervised process is not ready") +) + +// NativeOptions is the common policy input. Windows callers additionally get +// token/session selection and Job Object containment in the native build; the +// test adapter uses the same shell and output limits without OS handles. +type NativeOptions struct { + Shells ShellPaths + WorkRoot string + MaxWrapperBytes uint64 + MaxOutputChunk uint64 + WindowsTermGrace time.Duration + Now func() time.Time +} + +func (options NativeOptions) withDefaults() NativeOptions { + if options.MaxWrapperBytes == 0 { + options.MaxWrapperBytes = 10 << 20 + } + if options.MaxOutputChunk == 0 || options.MaxOutputChunk > 64<<10 { + options.MaxOutputChunk = 64 << 10 + } + if options.WindowsTermGrace <= 0 { + options.WindowsTermGrace = 10 * time.Second + } + if options.Now == nil { + options.Now = func() time.Time { return time.Now().UTC() } + } + return options +} + +type execSupervisor struct { + options NativeOptions + mu sync.Mutex + active map[domain.UUID]*execProcess +} + +func newExecSupervisor(options NativeOptions) *execSupervisor { + options = options.withDefaults() + return &execSupervisor{options: options, active: make(map[domain.UUID]*execProcess)} +} + +type execProcess struct { + issue domain.UUID + identity supervisor.EffectiveIdentity + cmd *exec.Cmd + stdin io.WriteCloser + stdout io.ReadCloser + stderr io.ReadCloser + outputs chan outputResult + done chan struct{} + started time.Time + waitFn func() (int32, bool, error) + killFn func(uint32) error + + mu sync.Mutex + finished bool + status supervisor.ExitStatus + waitErr error + closeIn sync.Once +} + +type outputResult struct { + chunk supervisor.OutputChunk + err error +} + +func (process *execProcess) IssueUUID() domain.UUID { return process.issue } + +func (process *execProcess) Identity() supervisor.EffectiveIdentity { return process.identity } + +func (process *execProcess) ReadOutput(ctx context.Context) (supervisor.OutputChunk, error) { + if process == nil || process.outputs == nil { + return supervisor.OutputChunk{}, ErrProcessNotReady + } + select { + case <-ctx.Done(): + return supervisor.OutputChunk{}, ctx.Err() + case result, ok := <-process.outputs: + if !ok { + return supervisor.OutputChunk{}, io.EOF + } + return result.chunk, result.err + } +} + +func (process *execProcess) Wait(ctx context.Context) (supervisor.ExitStatus, error) { + if process == nil || process.done == nil { + return supervisor.ExitStatus{}, ErrProcessNotReady + } + select { + case <-ctx.Done(): + return supervisor.ExitStatus{}, ctx.Err() + case <-process.done: + process.mu.Lock() + defer process.mu.Unlock() + return process.status, process.waitErr + } +} + +func (process *execProcess) WriteStdin(ctx context.Context, data []byte, appendNewline bool) error { + if process == nil || process.stdin == nil { + return ErrProcessNotReady + } + if appendNewline { + data = append(append([]byte(nil), data...), stdinLineEnding()...) + } + return writeWithContext(ctx, process.stdin, data) +} + +func (process *execProcess) CloseStdin(_ context.Context) error { + if process == nil || process.stdin == nil { + return ErrProcessNotReady + } + var err error + process.closeIn.Do(func() { err = process.stdin.Close() }) + return err +} + +func writeWithContext(ctx context.Context, writer io.Writer, data []byte) error { + if len(data) == 0 { + return nil + } + result := make(chan error, 1) + go func() { + _, err := writer.Write(data) + result <- err + }() + select { + case <-ctx.Done(): + return ctx.Err() + case err := <-result: + return err + } +} + +func (process *execProcess) setResult(status supervisor.ExitStatus, err error) { + process.mu.Lock() + process.status, process.waitErr, process.finished = status, err, true + process.mu.Unlock() + close(process.done) +} + +func (process *execProcess) startReaders(maxChunk uint64, remove func()) { + var readers sync.WaitGroup + read := func(stream rvboxv1.StreamKind, source io.ReadCloser) { + defer readers.Done() + defer source.Close() + limit := int(maxChunk) + if limit <= 0 { + limit = 64 << 10 + } + buffer := make([]byte, limit) + for { + count, err := source.Read(buffer) + if count > 0 { + data := append([]byte(nil), buffer[:count]...) + select { + case process.outputs <- outputResult{chunk: supervisor.OutputChunk{Stream: stream, Data: data}}: + default: + // The output channel is bounded. A reader must never hold a + // child pipe open while waiting for network I/O; dropping here + // is surfaced as an explicit read error to the caller. + select { + case process.outputs <- outputResult{err: fmt.Errorf("%w: output channel full", supervisor.ErrUnsupported)}: + case <-time.After(time.Second): + } + } + } + if err != nil { + if !errors.Is(err, io.EOF) { + select { + case process.outputs <- outputResult{err: err}: + default: + } + } + return + } + } + } + readers.Add(2) + go read(rvboxv1.StreamKind_STREAM_STDOUT, process.stdout) + go read(rvboxv1.StreamKind_STREAM_STDERR, process.stderr) + go func() { + readers.Wait() + close(process.outputs) + if remove != nil { + remove() + } + }() +} + +func materializeWrapper(directory string, wrapper Wrapper, now time.Time) (string, func(), error) { + if directory == "" || !now.IsZero() && now.Location() == nil { + return "", nil, ErrInvalidWorkingDirectory + } + if err := os.MkdirAll(directory, 0o700); err != nil { + return "", nil, err + } + temporary, err := os.CreateTemp(directory, ".rvbox-wrapper-*") + if err != nil { + return "", nil, err + } + temporaryName := temporary.Name() + cleanup := func() { + _ = temporary.Close() + _ = os.Remove(temporaryName) + } + if err := temporary.Chmod(0o600); err != nil { + cleanup() + return "", nil, err + } + if _, err := temporary.Write(wrapper.Bytes); err != nil { + cleanup() + return "", nil, err + } + if err := temporary.Sync(); err != nil { + cleanup() + return "", nil, err + } + if err := temporary.Close(); err != nil { + _ = os.Remove(temporaryName) + return "", nil, err + } + finalName := temporaryName + wrapper.Extension + if err := os.Rename(temporaryName, finalName); err != nil { + _ = os.Remove(temporaryName) + return "", nil, err + } + return finalName, func() { _ = os.Remove(finalName) }, nil +} + +func parseEnvironment(values []string) map[string]string { + result := make(map[string]string, len(values)) + for _, value := range values { + index := strings.IndexByte(value, '=') + if index <= 0 { + continue + } + result[value[:index]] = value[index+1:] + } + return result +} + +func (process *execProcess) terminate(code uint32) error { + if process == nil { + return ErrProcessNotReady + } + if process.killFn != nil { + return process.killFn(code) + } + if process.cmd == nil || process.cmd.Process == nil { + return ErrProcessNotReady + } + return process.cmd.Process.Kill() +} + +func (manager *execSupervisor) remove(issue domain.UUID, process *execProcess) { + manager.mu.Lock() + if manager.active[issue] == process { + delete(manager.active, issue) + } + manager.mu.Unlock() +} + +func validateExecutionSource(spec *supervisor.StartSpec) error { + if err := spec.Validate(); err != nil { + return err + } + if spec.Execution.GetScript() != nil && len(spec.ScriptBody) == 0 && spec.Execution.GetScript().GetSizeBytes() != 0 { + return ErrProcessNotReady + } + if len(spec.ExecutionProfiles) != 0 { + return fmt.Errorf("%w: execution profiles are not supported by this adapter yet", supervisor.ErrUnsupported) + } + return nil +} + +func sourceCommand(spec supervisor.StartSpec, wrapperPath string) (string, []string, error) { + source := spec.Execution.GetCommandText() + if wrapperPath != "" { + source = wrapperPath + } + switch spec.Execution.GetShellType() { + case rvboxv1.ShellType_SHELL_SH: + if wrapperPath != "" { + return "/bin/sh", []string{wrapperPath}, nil + } + return "/bin/sh", []string{"-c", source}, nil + case rvboxv1.ShellType_SHELL_BASH: + if wrapperPath != "" { + return "/bin/bash", []string{wrapperPath}, nil + } + return "/bin/bash", []string{"-c", source}, nil + default: + return "", nil, ErrUnsupportedShell + } +} + +// startPortable launches the same bounded pipe contract on Unix. It is kept +// private because v1 does not advertise a Unix client; tests use it to prove +// runner/supervisor ordering without a Windows host. +func (manager *execSupervisor) startPortable(ctx context.Context, spec supervisor.StartSpec) (supervisor.Process, error) { + if err := validateExecutionSource(&spec); err != nil { + return nil, err + } + manager.mu.Lock() + if _, exists := manager.active[spec.IssueUUID]; exists { + manager.mu.Unlock() + return nil, ErrProcessAlreadyRunning + } + manager.mu.Unlock() + var wrapperPath string + var cleanup func() + if spec.Execution.GetScript() != nil { + wrapper, err := BuildWrapper(spec.Execution, spec.ScriptBody, manager.options.MaxWrapperBytes) + if err != nil { + return nil, err + } + wrapperPath, cleanup, err = materializeWrapper(spec.WorkingDirectory, wrapper, manager.options.Now()) + if err != nil { + return nil, err + } + } + program, arguments, err := sourceCommand(spec, wrapperPath) + if err != nil { + if cleanup != nil { + cleanup() + } + return nil, err + } + return manager.startCommand(ctx, spec, program, arguments, supervisor.EffectiveIdentity{Context: "CURRENT_PROCESS", Elevated: false}, cleanup, nil) +} + +func (manager *execSupervisor) startCommand(ctx context.Context, spec supervisor.StartSpec, program string, arguments []string, identity supervisor.EffectiveIdentity, cleanup func(), configure func(*exec.Cmd) error) (supervisor.Process, error) { + command := exec.CommandContext(ctx, program, arguments...) + if configure != nil { + if err := configure(command); err != nil { + if cleanup != nil { + cleanup() + } + return nil, err + } + } + command.Dir = filepath.Clean(spec.WorkingDirectory) + base := parseEnvironment(os.Environ()) + entries, err := MergeEnvironment(base, spec.Environment) + if err != nil { + if cleanup != nil { + cleanup() + } + return nil, err + } + command.Env = make([]string, 0, len(entries)) + for _, entry := range entries { + command.Env = append(command.Env, entry.Key+"="+entry.Value) + } + stdinRead, stdin, err := os.Pipe() + if err != nil { + if cleanup != nil { + cleanup() + } + return nil, err + } + stdout, stdoutWrite, err := os.Pipe() + if err != nil { + _ = stdinRead.Close() + _ = stdin.Close() + if cleanup != nil { + cleanup() + } + return nil, err + } + stderr, stderrWrite, err := os.Pipe() + if err != nil { + _ = stdinRead.Close() + _ = stdin.Close() + _ = stdout.Close() + _ = stdoutWrite.Close() + if cleanup != nil { + cleanup() + } + return nil, err + } + command.Stdin = stdinRead + command.Stdout = stdoutWrite + command.Stderr = stderrWrite + if err := command.Start(); err != nil { + _ = stdinRead.Close() + _ = stdin.Close() + _ = stdout.Close() + _ = stdoutWrite.Close() + _ = stderr.Close() + _ = stderrWrite.Close() + if cleanup != nil { + cleanup() + } + return nil, err + } + _ = stdinRead.Close() + _ = stdoutWrite.Close() + _ = stderrWrite.Close() + started := manager.options.Now() + waitFn := func() (int32, bool, error) { + err := command.Wait() + var code int32 + if command.ProcessState != nil { + code = int32(command.ProcessState.ExitCode()) + } + return code, false, err + } + return manager.registerProcess(spec.IssueUUID, identity, command, stdin, stdout, stderr, started, waitFn, func(uint32) error { return command.Process.Kill() }, cleanup), nil +} + +func (manager *execSupervisor) registerProcess(issue domain.UUID, identity supervisor.EffectiveIdentity, command *exec.Cmd, stdin io.WriteCloser, stdout, stderr io.ReadCloser, started time.Time, waitFn func() (int32, bool, error), killFn func(uint32) error, cleanup func()) *execProcess { + process := &execProcess{issue: issue, identity: identity, cmd: command, stdin: stdin, stdout: stdout, stderr: stderr, outputs: make(chan outputResult, 32), done: make(chan struct{}), started: started, waitFn: waitFn, killFn: killFn} + manager.mu.Lock() + manager.active[issue] = process + manager.mu.Unlock() + process.startReaders(manager.options.MaxOutputChunk, cleanup) + go func() { + code, signaled, err := process.waitFn() + finished := manager.options.Now() + status := supervisor.ExitStatus{StartedAt: started, FinishedAt: finished, OutputDrained: false, Code: code, Signaled: signaled} + process.setResult(status, err) + manager.remove(issue, process) + }() + return process +} diff --git a/internal/client/supervisor/windows/native_other.go b/internal/client/supervisor/windows/native_other.go new file mode 100644 index 0000000..4bdf394 --- /dev/null +++ b/internal/client/supervisor/windows/native_other.go @@ -0,0 +1,80 @@ +//go:build !windows + +package windows + +import ( + "context" + "errors" + "os" + "time" + + "github.com/rvbox/rvbox/internal/client/supervisor" +) + +func stdinLineEnding() []byte { return []byte{'\n'} } + +// NewSupervisor returns the deterministic local adapter used by tests and by +// non-Windows development builds. The v1 client does not advertise this +// adapter as a supported Unix client; it exists so protocol/runtime tests can +// exercise real child processes without a Windows host. +func NewSupervisor(options NativeOptions) (supervisor.Supervisor, error) { + return newExecSupervisor(options), nil +} + +func (manager *execSupervisor) Start(ctx context.Context, spec supervisor.StartSpec) (supervisor.Process, error) { + return manager.startPortable(ctx, spec) +} + +func (manager *execSupervisor) Signal(ctx context.Context, process supervisor.Process, signal supervisor.SignalKind) (supervisor.SignalOutcome, error) { + if process == nil { + return supervisor.SignalOutcome{}, ErrProcessNotFound + } + executable, ok := process.(*execProcess) + if !ok { + return supervisor.SignalOutcome{}, ErrProcessNotFound + } + if signal != supervisor.SignalTerm && signal != supervisor.SignalKill { + return supervisor.SignalOutcome{}, errors.New("unsupported signal") + } + if signal == supervisor.SignalTerm { + if executable.cmd.Process != nil { + _ = executable.cmd.Process.Signal(os.Interrupt) + } + select { + case <-ctx.Done(): + return supervisor.SignalOutcome{}, ctx.Err() + case <-time.After(manager.options.WindowsTermGrace): + } + } + if err := executable.terminate(1); err != nil { + return supervisor.SignalOutcome{}, err + } + return supervisor.SignalOutcome{Delivered: true, Escalated: signal == supervisor.SignalTerm, Detail: "process terminated", ObservedAt: manager.options.Now()}, nil +} + +func (manager *execSupervisor) Snapshot(ctx context.Context, process supervisor.Process) (supervisor.ResourceSnapshot, error) { + if process == nil { + return supervisor.ResourceSnapshot{}, ErrProcessNotFound + } + select { + case <-ctx.Done(): + return supervisor.ResourceSnapshot{}, ctx.Err() + default: + } + return supervisor.ResourceSnapshot{ProcessCount: 1, ObservedAt: manager.options.Now(), Complete: false, Detail: "portable adapter does not expose aggregate process accounting"}, nil +} + +func (manager *execSupervisor) StopAll(_ context.Context) error { + manager.mu.Lock() + processes := make([]*execProcess, 0, len(manager.active)) + for _, process := range manager.active { + processes = append(processes, process) + } + manager.mu.Unlock() + for _, process := range processes { + if err := process.terminate(1); err != nil && !errors.Is(err, os.ErrProcessDone) { + return err + } + } + return nil +} diff --git a/internal/client/supervisor/windows/native_other_test.go b/internal/client/supervisor/windows/native_other_test.go new file mode 100644 index 0000000..498bf8c --- /dev/null +++ b/internal/client/supervisor/windows/native_other_test.go @@ -0,0 +1,93 @@ +//go:build !windows + +package windows + +import ( + "context" + "crypto/sha256" + "errors" + "io" + "strings" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/client/supervisor" + "github.com/rvbox/rvbox/internal/domain" +) + +func TestPortableSupervisorCapturesOutputAndSupportsStdin_HP_SUPERVISOR_03(t *testing.T) { + t.Parallel() + manager, err := NewSupervisor(NativeOptions{MaxOutputChunk: 8, WindowsTermGrace: 10 * time.Millisecond}) + if err != nil { + t.Fatal(err) + } + issue, err := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-0000000000c1") + if err != nil { + t.Fatal(err) + } + process, err := manager.Start(context.Background(), supervisor.StartSpec{IssueUUID: issue, CommandRevision: 1, WorkingDirectory: t.TempDir(), Execution: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_SH, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "read line; printf 'out:%s\\n' \"$line\"; printf 'err\\n' >&2"}}}) + if err != nil { + t.Fatal(err) + } + if err := process.WriteStdin(context.Background(), []byte("hello"), true); err != nil { + t.Fatal(err) + } + if err := process.CloseStdin(context.Background()); err != nil { + t.Fatal(err) + } + status, err := process.Wait(context.Background()) + if err != nil || status.Code != 0 { + t.Fatalf("wait = %+v, %v", status, err) + } + var output strings.Builder + streams := map[rvboxv1.StreamKind]bool{} + for { + chunk, err := process.ReadOutput(context.Background()) + if errors.Is(err, io.EOF) { + break + } + if err != nil { + t.Fatal(err) + } + output.Write(chunk.Data) + streams[chunk.Stream] = true + } + if !strings.Contains(output.String(), "out:hello\n") || !strings.Contains(output.String(), "err\n") || !streams[rvboxv1.StreamKind_STREAM_STDOUT] || !streams[rvboxv1.StreamKind_STREAM_STDERR] { + t.Fatalf("captured output = %q streams=%v", output.String(), streams) + } + second, err := manager.Start(context.Background(), supervisor.StartSpec{IssueUUID: issue, CommandRevision: 1, WorkingDirectory: t.TempDir(), Execution: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_SH, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "true"}}}) + if err != nil { + t.Fatalf("restart after terminal = %v", err) + } + _, _ = second.Wait(context.Background()) +} + +func TestPortableSupervisorScriptMaterializationAndCleanup_HP_SUPERVISOR_04(t *testing.T) { + t.Parallel() + manager, err := NewSupervisor(NativeOptions{MaxWrapperBytes: 1024}) + if err != nil { + t.Fatal(err) + } + issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-0000000000c2") + body := []byte("printf script-ok") + process, err := manager.Start(context.Background(), supervisor.StartSpec{IssueUUID: issue, CommandRevision: 1, WorkingDirectory: t.TempDir(), ScriptBody: body, Execution: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_SH, Source: &rvboxv1.ExecutionSpec_Script{Script: &rvboxv1.ScriptDescriptor{SizeBytes: uint64(len(body)), Sha256: digest(body)}}}}) + if err != nil { + t.Fatal(err) + } + if _, err := process.Wait(context.Background()); err != nil { + t.Fatal(err) + } + chunk, err := process.ReadOutput(context.Background()) + if err != nil || string(chunk.Data) != "script-ok" { + t.Fatalf("script output = %+v, %v", chunk, err) + } + if _, err := process.ReadOutput(context.Background()); !errors.Is(err, io.EOF) { + t.Fatalf("script output close = %v", err) + } +} + +func digest(value []byte) []byte { + result := sha256.Sum256(value) + return result[:] +} diff --git a/internal/client/supervisor/windows/native_windows.go b/internal/client/supervisor/windows/native_windows.go new file mode 100644 index 0000000..631b862 --- /dev/null +++ b/internal/client/supervisor/windows/native_windows.go @@ -0,0 +1,476 @@ +//go:build windows + +package windows + +// The Windows adapter deliberately keeps all Win32 handles in this file. The +// selector in selection.go is pure policy; this layer obtains one verified +// primary token, creates a suspended child with an explicit handle list, and +// puts it in a kill-on-close Job before releasing it. + +import ( + "context" + "errors" + "fmt" + "os" + "os/exec" + "sync" + "syscall" + "time" + "unsafe" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/client/supervisor" + winapi "golang.org/x/sys/windows" +) + +const ( + logon32LogonService = 5 + logon32ProviderDefault = 0 + securitySystemRID = "S-1-5-18" +) + +var ( + advapi32 = syscall.NewLazyDLL("advapi32.dll") + procLogonUserW = advapi32.NewProc("LogonUserW") +) + +func stdinLineEnding() []byte { return []byte{'\r', '\n'} } + +type nativeHandles struct { + process winapi.Handle + job winapi.Handle + pid uint32 + close sync.Once +} + +// NewSupervisor constructs the machine-wide Windows implementation. The +// service process is expected to run as LocalSystem; token selection verifies +// that assumption when a command is started and records the selected context. +func NewSupervisor(options NativeOptions) (supervisor.Supervisor, error) { + options = options.withDefaults() + return newExecSupervisor(options), nil +} + +func (manager *execSupervisor) Start(ctx context.Context, spec supervisor.StartSpec) (supervisor.Process, error) { + if err := validateExecutionSource(&spec); err != nil { + return nil, err + } + if spec.Execution.GetShellType() != rvboxv1.ShellType_SHELL_CMD && spec.Execution.GetShellType() != rvboxv1.ShellType_SHELL_POWERSHELL { + return nil, ErrUnsupportedShell + } + manager.mu.Lock() + if _, exists := manager.active[spec.IssueUUID]; exists { + manager.mu.Unlock() + return nil, ErrProcessAlreadyRunning + } + manager.mu.Unlock() + + plan, err := ResolveShell(spec.Execution.GetShellType(), manager.options.Shells) + if err != nil { + return nil, err + } + if err := verifyExecutable(plan.ApplicationName); err != nil { + return nil, err + } + wrapper, err := BuildWrapper(spec.Execution, spec.ScriptBody, manager.options.MaxWrapperBytes) + if err != nil { + return nil, err + } + wrapperPath, cleanup, err := materializeWrapper(spec.WorkingDirectory, wrapper, manager.options.Now()) + if err != nil { + return nil, err + } + fail := func(cause error) (supervisor.Process, error) { + if cleanup != nil { + cleanup() + } + return nil, cause + } + + token, identity, err := manager.selectToken(spec.Execution.GetElevated()) + if err != nil { + return fail(err) + } + defer token.Close() + baseEnvironment, err := token.Environ(false) + if err != nil { + return fail(fmt.Errorf("build token environment: %w", err)) + } + environment, err := BuildEnvironmentBlock(parseEnvironment(baseEnvironment), spec.Environment) + if err != nil { + return fail(err) + } + launch, err := plan.BuildLaunchPlan(wrapperPath, spec.WorkingDirectory, environment) + if err != nil { + return fail(err) + } + + stdinRead, stdinWrite, stdoutRead, stdoutWrite, stderrRead, stderrWrite, err := createStandardPipes() + if err != nil { + return fail(err) + } + closeFiles := func() { + for _, file := range []*os.File{stdinRead, stdinWrite, stdoutRead, stdoutWrite, stderrRead, stderrWrite} { + if file != nil { + _ = file.Close() + } + } + } + pipesTransferred := false + defer func() { + // Parent-side handles are retained only after successful process + // creation. Any error path closes both ends here. + if !pipesTransferred { + closeFiles() + } + }() + + job, err := createKillOnCloseJob() + if err != nil { + closeFiles() + return fail(fmt.Errorf("create command Job: %w", err)) + } + cleanupJob := true + defer func() { + if cleanupJob { + _ = winapi.CloseHandle(job) + } + }() + + application, err := winapi.UTF16PtrFromString(launch.ApplicationName) + if err != nil { + return fail(err) + } + commandLine, err := winapi.UTF16FromString(launch.CommandLine) + if err != nil { + return fail(err) + } + workingDirectory, err := winapi.UTF16PtrFromString(launch.WorkingDirectory) + if err != nil { + return fail(err) + } + attributeList, err := winapi.NewProcThreadAttributeList(1) + if err != nil { + return fail(err) + } + defer attributeList.Delete() + childHandles := []winapi.Handle{winapi.Handle(stdinRead.Fd()), winapi.Handle(stdoutWrite.Fd()), winapi.Handle(stderrWrite.Fd())} + if err := attributeList.Update(winapi.PROC_THREAD_ATTRIBUTE_HANDLE_LIST, unsafe.Pointer(&childHandles[0]), uintptr(len(childHandles))*unsafe.Sizeof(childHandles[0])); err != nil { + return fail(err) + } + startup := winapi.StartupInfoEx{} + startup.Cb = uint32(unsafe.Sizeof(startup)) + startup.Flags = winapi.STARTF_USESTDHANDLES | winapi.STARTF_USESHOWWINDOW + startup.ShowWindow = winapi.SW_HIDE + startup.StdInput = childHandles[0] + startup.StdOutput = childHandles[1] + startup.StdErr = childHandles[2] + startup.ProcThreadAttributeList = attributeList.List() + var processInfo winapi.ProcessInformation + flags := uint32(winapi.CREATE_NEW_CONSOLE | winapi.CREATE_SUSPENDED | winapi.CREATE_UNICODE_ENVIRONMENT | winapi.EXTENDED_STARTUPINFO_PRESENT) + var environmentPointer *uint16 + if len(environment) > 0 { + environmentPointer = &environment[0] + } + if err := winapi.CreateProcessAsUser(token, application, &commandLine[0], nil, nil, true, flags, environmentPointer, workingDirectory, &startup.StartupInfo, &processInfo); err != nil { + return fail(fmt.Errorf("create suspended command process: %w", err)) + } + // The child owns these handles after CreateProcessAsUser returns. Keep only + // the three parent ends and the process/job handles in the daemon. + _ = stdinRead.Close() + _ = stdoutWrite.Close() + _ = stderrWrite.Close() + if err := winapi.AssignProcessToJobObject(job, processInfo.Process); err != nil { + _ = winapi.TerminateProcess(processInfo.Process, 1) + _ = winapi.CloseHandle(processInfo.Process) + _ = winapi.CloseHandle(processInfo.Thread) + return fail(fmt.Errorf("assign command to Job: %w", err)) + } + if _, err := winapi.ResumeThread(processInfo.Thread); err != nil { + _ = winapi.TerminateJobObject(job, 1) + _ = winapi.CloseHandle(processInfo.Process) + _ = winapi.CloseHandle(processInfo.Thread) + return fail(fmt.Errorf("release suspended command: %w", err)) + } + pipesTransferred = true + _ = winapi.CloseHandle(processInfo.Thread) + started := manager.options.Now() + handles := &nativeHandles{process: processInfo.Process, job: job, pid: processInfo.ProcessId} + cleanupJob = false + command := &exec.Cmd{Process: osProcess(processInfo.ProcessId)} + waitFn := func() (int32, bool, error) { + _, waitErr := winapi.WaitForSingleObject(processInfo.Process, winapi.INFINITE) + var code uint32 + if err := winapi.GetExitCodeProcess(processInfo.Process, &code); err != nil && waitErr == nil { + waitErr = err + } + handles.close.Do(func() { + _ = winapi.CloseHandle(processInfo.Process) + _ = winapi.CloseHandle(job) + }) + return int32(code), false, waitErr + } + killFn := func(code uint32) error { + return winapi.TerminateJobObject(job, code) + } + process := manager.registerProcess(spec.IssueUUID, identity, command, stdinWrite, stdoutRead, stderrRead, started, waitFn, killFn, cleanup) + return process, nil +} + +func osProcess(pid uint32) *os.Process { + process, err := os.FindProcess(int(pid)) + if err != nil { + return &os.Process{} + } + return process +} + +func verifyExecutable(path string) error { + info, err := os.Stat(path) + if err != nil { + return fmt.Errorf("stat configured shell %q: %w", path, err) + } + if !info.Mode().IsRegular() { + return fmt.Errorf("configured shell %q is not a regular file", path) + } + return nil +} + +func createStandardPipes() (*os.File, *os.File, *os.File, *os.File, *os.File, *os.File, error) { + security := &winapi.SecurityAttributes{Length: uint32(unsafe.Sizeof(winapi.SecurityAttributes{})), InheritHandle: 1} + var stdinReadHandle, stdinWriteHandle winapi.Handle + var stdoutReadHandle, stdoutWriteHandle winapi.Handle + var stderrReadHandle, stderrWriteHandle winapi.Handle + if err := winapi.CreatePipe(&stdinReadHandle, &stdinWriteHandle, security, 0); err != nil { + return nil, nil, nil, nil, nil, nil, err + } + if err := winapi.CreatePipe(&stdoutReadHandle, &stdoutWriteHandle, security, 0); err != nil { + _ = winapi.CloseHandle(stdinReadHandle) + _ = winapi.CloseHandle(stdinWriteHandle) + return nil, nil, nil, nil, nil, nil, err + } + if err := winapi.CreatePipe(&stderrReadHandle, &stderrWriteHandle, security, 0); err != nil { + for _, handle := range []winapi.Handle{stdinReadHandle, stdinWriteHandle, stdoutReadHandle, stdoutWriteHandle} { + _ = winapi.CloseHandle(handle) + } + return nil, nil, nil, nil, nil, nil, err + } + for _, handle := range []winapi.Handle{stdinWriteHandle, stdoutReadHandle, stderrReadHandle} { + if err := winapi.SetHandleInformation(handle, winapi.HANDLE_FLAG_INHERIT, 0); err != nil { + for _, closeHandle := range []winapi.Handle{stdinReadHandle, stdinWriteHandle, stdoutReadHandle, stdoutWriteHandle, stderrReadHandle, stderrWriteHandle} { + _ = winapi.CloseHandle(closeHandle) + } + return nil, nil, nil, nil, nil, nil, err + } + } + return os.NewFile(uintptr(stdinReadHandle), "rvbox-stdin-read"), os.NewFile(uintptr(stdinWriteHandle), "rvbox-stdin-write"), os.NewFile(uintptr(stdoutReadHandle), "rvbox-stdout-read"), os.NewFile(uintptr(stdoutWriteHandle), "rvbox-stdout-write"), os.NewFile(uintptr(stderrReadHandle), "rvbox-stderr-read"), os.NewFile(uintptr(stderrWriteHandle), "rvbox-stderr-write"), nil +} + +func createKillOnCloseJob() (winapi.Handle, error) { + job, err := winapi.CreateJobObject(nil, nil) + if err != nil { + return 0, err + } + info := winapi.JOBOBJECT_EXTENDED_LIMIT_INFORMATION{} + info.BasicLimitInformation.LimitFlags = winapi.JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE + if _, err := winapi.SetInformationJobObject(job, winapi.JobObjectExtendedLimitInformation, uintptr(unsafe.Pointer(&info)), uint32(unsafe.Sizeof(info))); err != nil { + _ = winapi.CloseHandle(job) + return 0, err + } + return job, nil +} + +func (manager *execSupervisor) selectToken(elevated bool) (winapi.Token, supervisor.EffectiveIdentity, error) { + candidates, err := DiscoverActiveSessions() + if err != nil { + // A failed WTS enumeration is treated as no usable interactive + // session; LocalSystem still provides a deterministic service path. + candidates = nil + } + selection := Select(SelectionInput{Elevated: elevated, ActiveSessions: candidates, ActiveSystemAvailable: true, LocalServiceAvailable: true, LocalSystemAvailable: true}) + if selection.Effective == nil { + if selection.Error != nil { + return 0, supervisor.EffectiveIdentity{}, selection.Error + } + return 0, supervisor.EffectiveIdentity{}, errors.New("Windows execution context selection failed") + } + var selected *SessionCandidate + if selection.Effective.SessionID != nil { + for index := range candidates { + if candidates[index].SessionID == *selection.Effective.SessionID { + selected = &candidates[index] + break + } + } + } + for _, attempt := range selection.Attempts { + token, identity, err := openTokenForAttempt(attempt.Context, selected) + if err == nil { + return token, identity, nil + } + if !elevated { + return 0, supervisor.EffectiveIdentity{}, err + } + } + // The pure selector stops as soon as ACTIVE_SYSTEM is available. A native + // privilege/session operation can still fail (for example, SeTcb was + // removed), so the final LOCAL_SYSTEM fallback is attempted here before + // launch preparation, never by retrying a created process. + if elevated { + if token, identity, err := openTokenForAttempt(ContextLocalSystem, nil); err == nil { + return token, identity, nil + } + } + return 0, supervisor.EffectiveIdentity{}, errors.New("all Windows execution contexts failed before launch preparation") +} + +func openTokenForAttempt(contextName ExecutionContext, candidate *SessionCandidate) (winapi.Token, supervisor.EffectiveIdentity, error) { + switch contextName { + case ContextActiveUser, ContextActiveUserElevated, ContextActiveSystem: + if candidate == nil { + return 0, supervisor.EffectiveIdentity{}, errors.New("active execution context has no selected session") + } + var token winapi.Token + if err := winapi.WTSQueryUserToken(candidate.SessionID, &token); err != nil { + return 0, supervisor.EffectiveIdentity{}, err + } + if contextName == ContextActiveUserElevated { + if !token.IsElevated() { + linked, err := token.GetLinkedToken() + _ = token.Close() + if err != nil { + return 0, supervisor.EffectiveIdentity{}, err + } + token = linked + } + } else if contextName == ContextActiveSystem { + _ = token.Close() + serviceToken, identity, err := duplicateServiceTokenForSession(candidate.SessionID) + if err != nil { + return 0, supervisor.EffectiveIdentity{}, err + } + identity.Context = string(ContextActiveSystem) + return serviceToken, identity, nil + } else if token.IsElevated() { + _ = token.Close() + return 0, supervisor.EffectiveIdentity{}, errors.New("active-user token is elevated and no restricted medium token was available") + } + identity := supervisor.EffectiveIdentity{Context: string(contextName), SessionID: candidate.SessionID, UserSID: candidate.UserSID, Elevated: contextName != ContextActiveUser, Integrity: map[bool]string{true: "high", false: "medium"}[contextName != ContextActiveUser]} + return token, identity, nil + case ContextLocalService: + token, err := logonLocalService() + if err != nil { + return 0, supervisor.EffectiveIdentity{}, err + } + return token, supervisor.EffectiveIdentity{Context: string(ContextLocalService), Elevated: false, Integrity: "medium"}, nil + case ContextLocalSystem: + return duplicateServiceToken() + default: + return 0, supervisor.EffectiveIdentity{}, errors.New("unknown Windows execution context") + } +} + +func duplicateServiceToken() (winapi.Token, supervisor.EffectiveIdentity, error) { + return duplicateServiceTokenForSession(0) +} + +func duplicateServiceTokenForSession(sessionID uint32) (winapi.Token, supervisor.EffectiveIdentity, error) { + var source winapi.Token + if err := winapi.OpenProcessToken(winapi.CurrentProcess(), winapi.TOKEN_ALL_ACCESS, &source); err != nil { + return 0, supervisor.EffectiveIdentity{}, err + } + defer source.Close() + var target winapi.Token + if err := winapi.DuplicateTokenEx(source, winapi.TOKEN_ALL_ACCESS, nil, winapi.SecurityImpersonation, winapi.TokenPrimary, &target); err != nil { + return 0, supervisor.EffectiveIdentity{}, err + } + if sessionID != 0 { + if err := winapi.SetTokenInformation(target, winapi.TokenSessionId, (*byte)(unsafe.Pointer(&sessionID)), uint32(unsafe.Sizeof(sessionID))); err != nil { + _ = target.Close() + return 0, supervisor.EffectiveIdentity{}, err + } + } + user, err := target.GetTokenUser() + if err != nil || user.User.Sid == nil || user.User.Sid.String() != securitySystemRID { + _ = target.Close() + if err != nil { + return 0, supervisor.EffectiveIdentity{}, err + } + return 0, supervisor.EffectiveIdentity{}, errors.New("duplicated service token is not LocalSystem") + } + return target, supervisor.EffectiveIdentity{Context: string(ContextLocalSystem), SessionID: sessionID, UserSID: user.User.Sid.String(), Elevated: true, Integrity: "system"}, nil +} + +func logonLocalService() (winapi.Token, error) { + account, _ := syscall.UTF16PtrFromString("LocalService") + domainName, _ := syscall.UTF16PtrFromString("NT AUTHORITY") + var token winapi.Token + r, _, callErr := procLogonUserW.Call(uintptr(unsafe.Pointer(account)), uintptr(unsafe.Pointer(domainName)), 0, logon32LogonService, logon32ProviderDefault, uintptr(unsafe.Pointer(&token))) + if r == 0 { + if callErr != syscall.Errno(0) { + return 0, callErr + } + return 0, syscall.GetLastError() + } + return token, nil +} + +func (manager *execSupervisor) Signal(ctx context.Context, process supervisor.Process, signal supervisor.SignalKind) (supervisor.SignalOutcome, error) { + if process == nil { + return supervisor.SignalOutcome{}, ErrProcessNotFound + } + native, ok := process.(*execProcess) + if !ok || native.killFn == nil { + return supervisor.SignalOutcome{}, ErrProcessNotFound + } + if signal != supervisor.SignalTerm && signal != supervisor.SignalKill { + return supervisor.SignalOutcome{}, errors.New("unsupported signal") + } + if signal == supervisor.SignalTerm { + // The command has its own hidden console. A full AttachConsole/control + // helper is intentionally isolated from the Job kill path; if it is not + // available, the bounded grace period ends in an explicit Job kill. + select { + case <-ctx.Done(): + return supervisor.SignalOutcome{}, ctx.Err() + case <-time.After(manager.options.WindowsTermGrace): + } + } + if err := native.killFn(1); err != nil { + return supervisor.SignalOutcome{}, err + } + return supervisor.SignalOutcome{Delivered: true, Escalated: signal == supervisor.SignalTerm, Detail: "Windows Job terminated", ObservedAt: manager.options.Now()}, nil +} + +func (manager *execSupervisor) Snapshot(ctx context.Context, process supervisor.Process) (supervisor.ResourceSnapshot, error) { + if process == nil { + return supervisor.ResourceSnapshot{}, ErrProcessNotFound + } + native, ok := process.(*execProcess) + if !ok || native.cmd == nil || native.cmd.Process == nil { + return supervisor.ResourceSnapshot{}, ErrProcessNotFound + } + select { + case <-ctx.Done(): + return supervisor.ResourceSnapshot{}, ctx.Err() + default: + } + return supervisor.ResourceSnapshot{ProcessCount: 1, ObservedAt: manager.options.Now(), Complete: false, Detail: "Windows Job accounting is available after native completion integration"}, nil +} + +func (manager *execSupervisor) StopAll(_ context.Context) error { + manager.mu.Lock() + processes := make([]*execProcess, 0, len(manager.active)) + for _, process := range manager.active { + processes = append(processes, process) + } + manager.mu.Unlock() + for _, process := range processes { + if process.killFn != nil { + if err := process.killFn(1); err != nil && !errors.Is(err, winapi.ERROR_INVALID_HANDLE) { + return err + } + } + } + return nil +} diff --git a/internal/client/windowsservice/service.go b/internal/client/windowsservice/service.go new file mode 100644 index 0000000..fcc2053 --- /dev/null +++ b/internal/client/windowsservice/service.go @@ -0,0 +1,124 @@ +// Package windowsservice owns the machine-wide Windows service contract. The +// policy and transition model are platform-neutral so they can be exercised +// on Linux; service manager calls live in build-tagged adapters. +package windowsservice + +import ( + "context" + "errors" + "fmt" + "strings" + + winlaunch "github.com/rvbox/rvbox/internal/client/supervisor/windows" +) + +const ( + Name = "RVBoxClient" + DisplayName = "RVBox Client" + Description = "RVBox Windows client daemon and command supervisor" +) + +var ( + ErrInvalidInstallSpec = errors.New("invalid Windows service install specification") + ErrUnsupported = errors.New("Windows service management is unavailable on this platform") + ErrInvalidTransition = errors.New("invalid Windows service lifecycle transition") +) + +type StartupMode string + +const ( + StartupAutomatic StartupMode = "automatic" + StartupManual StartupMode = "manual" +) + +// InstallSpec is the complete immutable service image contract. The service +// is always one machine-wide LocalSystem process; config selection is explicit +// and never delegated to PATH, the current directory, or Task Scheduler. +type InstallSpec struct { + ExecutablePath string + ConfigPath string + Startup StartupMode +} + +func (spec InstallSpec) Validate() error { + if !winlaunch.ValidAbsoluteWindowsPath(spec.ExecutablePath) || !winlaunch.ValidAbsoluteWindowsPath(spec.ConfigPath) { + return ErrInvalidInstallSpec + } + if strings.ContainsRune(spec.ExecutablePath, 0) || strings.ContainsRune(spec.ConfigPath, 0) { + return ErrInvalidInstallSpec + } + if spec.Startup != StartupAutomatic && spec.Startup != StartupManual { + return ErrInvalidInstallSpec + } + return nil +} + +func (spec InstallSpec) String() string { + return fmt.Sprintf("%s config=%s startup=%s", spec.ExecutablePath, spec.ConfigPath, spec.Startup) +} + +type State uint8 + +const ( + StateUnknown State = iota + StateStopped + StateStartPending + StateRunning + StateStopPending + StatePaused +) + +type Command uint8 + +const ( + CommandStart Command = iota + 1 + CommandStop + CommandRestart +) + +// Transition validates only state changes that the service control adapter is +// allowed to request. Pending states are intentionally not collapsed into a +// successful result: callers must poll SCM and observe the final state. +func Transition(state State, command Command) (State, error) { + switch command { + case CommandStart: + switch state { + case StateStopped: + return StateStartPending, nil + case StateRunning, StateStartPending: + return state, nil + default: + return StateUnknown, ErrInvalidTransition + } + case CommandStop: + switch state { + case StateRunning, StatePaused: + return StateStopPending, nil + case StateStopped, StateStopPending: + return state, nil + default: + return StateUnknown, ErrInvalidTransition + } + case CommandRestart: + switch state { + case StateStopped: + return StateStartPending, nil + case StateRunning, StatePaused: + return StateStopPending, nil + case StateStartPending, StateStopPending: + return state, nil + default: + return StateUnknown, ErrInvalidTransition + } + default: + return StateUnknown, ErrInvalidTransition + } +} + +// Install, Uninstall, Start, and Stop are intentionally narrow. Their +// platform implementations return ErrUnsupported on non-Windows builds. +func Install(spec InstallSpec) error { return installNative(spec) } +func Uninstall() error { return uninstallNative() } +func Start() error { return startNative() } +func Stop(timeoutSeconds uint32) error { return stopNative(timeoutSeconds) } +func Run(run func(context.Context) error) error { return runNative(run) } diff --git a/internal/client/windowsservice/service_other.go b/internal/client/windowsservice/service_other.go new file mode 100644 index 0000000..d43ef20 --- /dev/null +++ b/internal/client/windowsservice/service_other.go @@ -0,0 +1,11 @@ +//go:build !windows + +package windowsservice + +import "context" + +func installNative(InstallSpec) error { return ErrUnsupported } +func uninstallNative() error { return ErrUnsupported } +func startNative() error { return ErrUnsupported } +func stopNative(uint32) error { return ErrUnsupported } +func runNative(func(context.Context) error) error { return ErrUnsupported } diff --git a/internal/client/windowsservice/service_test.go b/internal/client/windowsservice/service_test.go new file mode 100644 index 0000000..82cc05a --- /dev/null +++ b/internal/client/windowsservice/service_test.go @@ -0,0 +1,59 @@ +package windowsservice + +import ( + "errors" + "testing" +) + +func TestInstallSpecValidation_HP_WINSVC_01(t *testing.T) { + t.Parallel() + valid := InstallSpec{ExecutablePath: `C:\Program Files\RVBox\rvbox.exe`, ConfigPath: `C:\ProgramData\RVBox\client.toml`, Startup: StartupAutomatic} + if err := valid.Validate(); err != nil { + t.Fatalf("valid install spec rejected: %v", err) + } + for _, invalid := range []InstallSpec{ + {ExecutablePath: `rvbox.exe`, ConfigPath: valid.ConfigPath, Startup: StartupAutomatic}, + {ExecutablePath: valid.ExecutablePath, ConfigPath: `client.toml`, Startup: StartupAutomatic}, + {ExecutablePath: valid.ExecutablePath, ConfigPath: valid.ConfigPath, Startup: StartupMode("disabled")}, + } { + if !errors.Is(invalid.Validate(), ErrInvalidInstallSpec) { + t.Fatalf("invalid install spec accepted: %#v", invalid) + } + } +} + +func TestServiceTransitionsAreIdempotentAndExplicit_HP_WINSVC_02(t *testing.T) { + t.Parallel() + tests := []struct { + state State + command Command + want State + }{ + {StateStopped, CommandStart, StateStartPending}, + {StateStartPending, CommandStart, StateStartPending}, + {StateRunning, CommandStop, StateStopPending}, + {StateStopPending, CommandStop, StateStopPending}, + {StateStopped, CommandRestart, StateStartPending}, + {StateRunning, CommandRestart, StateStopPending}, + } + for _, test := range tests { + got, err := Transition(test.state, test.command) + if err != nil || got != test.want { + t.Errorf("Transition(%v,%v) = %v,%v; want %v,nil", test.state, test.command, got, err, test.want) + } + } + for _, test := range []struct { + state State + command Command + }{{StateUnknown, CommandStart}, {StateStartPending, CommandStop}} { + if test.state == StateStopped && test.command == CommandStop { + // Stopping an already stopped service is deliberately idempotent. + continue + } + if _, err := Transition(test.state, test.command); test.state == StateStopped && test.command == CommandStop { + t.Fatalf("unreachable idempotent case returned error: %v", err) + } else if err == nil { + t.Errorf("Transition(%v,%v) unexpectedly accepted", test.state, test.command) + } + } +} diff --git a/internal/client/windowsservice/service_windows.go b/internal/client/windowsservice/service_windows.go new file mode 100644 index 0000000..a4234dd --- /dev/null +++ b/internal/client/windowsservice/service_windows.go @@ -0,0 +1,207 @@ +//go:build windows + +package windowsservice + +import ( + "context" + "errors" + "fmt" + "syscall" + "time" + + "golang.org/x/sys/windows" + "golang.org/x/sys/windows/svc" + "golang.org/x/sys/windows/svc/mgr" +) + +func connect() (*mgr.Mgr, error) { + return mgr.Connect() +} + +func installNative(spec InstallSpec) error { + if err := spec.Validate(); err != nil { + return err + } + manager, err := connect() + if err != nil { + return fmt.Errorf("connect to service control manager: %w", err) + } + defer manager.Disconnect() + startup := uint32(mgr.StartAutomatic) + if spec.Startup == StartupManual { + startup = mgr.StartManual + } + configuration := mgr.Config{ + ServiceType: windows.SERVICE_WIN32_OWN_PROCESS, + StartType: startup, + ErrorControl: mgr.ErrorNormal, + DisplayName: DisplayName, + Description: Description, + ServiceStartName: "LocalSystem", + } + service, openErr := manager.OpenService(Name) + if errors.Is(openErr, windows.ERROR_SERVICE_DOES_NOT_EXIST) { + service, err = manager.CreateService(Name, spec.ExecutablePath, configuration, "--service", "--config", spec.ConfigPath) + if err != nil { + return fmt.Errorf("create %s service: %w", Name, err) + } + } else if openErr != nil { + return fmt.Errorf("open %s service: %w", Name, openErr) + } else { + configuration, err = service.Config() + if err != nil { + return fmt.Errorf("query %s service configuration: %w", Name, err) + } + configuration.ServiceType = windows.SERVICE_WIN32_OWN_PROCESS + configuration.StartType = startup + configuration.ErrorControl = mgr.ErrorNormal + configuration.BinaryPathName = serviceImage(spec) + configuration.DisplayName = DisplayName + configuration.Description = Description + configuration.ServiceStartName = "LocalSystem" + if err := service.UpdateConfig(configuration); err != nil { + return fmt.Errorf("update %s service configuration: %w", Name, err) + } + } + defer service.Close() + if err := service.Start(); err != nil && !errors.Is(err, windows.ERROR_SERVICE_ALREADY_RUNNING) { + return fmt.Errorf("start %s service: %w", Name, err) + } + return nil +} + +func serviceImage(spec InstallSpec) string { + return syscall.EscapeArg(spec.ExecutablePath) + " --service --config " + syscall.EscapeArg(spec.ConfigPath) +} + +func uninstallNative() error { + manager, err := connect() + if err != nil { + return fmt.Errorf("connect to service control manager: %w", err) + } + defer manager.Disconnect() + service, err := manager.OpenService(Name) + if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) { + return nil + } + if err != nil { + return fmt.Errorf("open %s service: %w", Name, err) + } + defer service.Close() + if err := stopService(service, 30*time.Second); err != nil { + return err + } + if err := service.Delete(); err != nil && !errors.Is(err, windows.ERROR_SERVICE_MARKED_FOR_DELETE) { + return fmt.Errorf("delete %s service: %w", Name, err) + } + return nil +} + +func startNative() error { + manager, err := connect() + if err != nil { + return err + } + defer manager.Disconnect() + service, err := manager.OpenService(Name) + if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) { + return fmt.Errorf("%s service is not installed", Name) + } + if err != nil { + return err + } + defer service.Close() + if err := service.Start(); err != nil && !errors.Is(err, windows.ERROR_SERVICE_ALREADY_RUNNING) { + return err + } + return nil +} + +func stopNative(timeoutSeconds uint32) error { + manager, err := connect() + if err != nil { + return err + } + defer manager.Disconnect() + service, err := manager.OpenService(Name) + if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) { + return nil + } + if err != nil { + return err + } + defer service.Close() + timeout := 30 * time.Second + if timeoutSeconds > 0 { + timeout = time.Duration(timeoutSeconds) * time.Second + } + return stopService(service, timeout) +} + +func stopService(service *mgr.Service, timeout time.Duration) error { + status, err := service.Query() + if err != nil { + return fmt.Errorf("query %s service: %w", Name, err) + } + if status.State == svc.Stopped { + return nil + } + if _, err := service.Control(svc.Stop); err != nil && !errors.Is(err, windows.ERROR_SERVICE_NOT_ACTIVE) { + return fmt.Errorf("stop %s service: %w", Name, err) + } + deadline := time.Now().Add(timeout) + for { + status, err = service.Query() + if err != nil { + return fmt.Errorf("query %s service while stopping: %w", Name, err) + } + if status.State == svc.Stopped { + return nil + } + if time.Now().After(deadline) { + return fmt.Errorf("timed out stopping %s service", Name) + } + time.Sleep(100 * time.Millisecond) + } +} + +type handler struct { + run func(context.Context) error +} + +func (serviceHandler handler) Execute(_ []string, changes <-chan svc.ChangeRequest, status chan<- svc.Status) (bool, uint32) { + if serviceHandler.run == nil { + return false, 1 + } + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + status <- svc.Status{State: svc.StartPending, WaitHint: 10_000} + done := make(chan error, 1) + go func() { done <- serviceHandler.run(ctx) }() + status <- svc.Status{State: svc.Running, Accepts: svc.AcceptStop | svc.AcceptShutdown} + for { + select { + case request := <-changes: + if request.Cmd == svc.Stop || request.Cmd == svc.Shutdown { + status <- svc.Status{State: svc.StopPending, WaitHint: 30_000} + cancel() + err := <-done + if err != nil { + return true, 1 + } + status <- svc.Status{State: svc.Stopped} + return false, 0 + } + case err := <-done: + if err != nil { + return true, 1 + } + status <- svc.Status{State: svc.Stopped} + return false, 0 + } + } +} + +func runNative(run func(context.Context) error) error { + return svc.Run(Name, handler{run: run}) +} diff --git a/internal/client/windowstray/protocol.go b/internal/client/windowstray/protocol.go new file mode 100644 index 0000000..8bb99a1 --- /dev/null +++ b/internal/client/windowstray/protocol.go @@ -0,0 +1,126 @@ +// Package windowstray contains the small, versioned protocol between the +// per-session notification-area process and the machine-wide service. The +// tray never receives command payloads or opens the client spool. +package windowstray + +import ( + "bytes" + "encoding/binary" + "errors" + "fmt" + "unicode/utf8" +) + +const ( + protocolVersion uint16 = 1 + maxFrameBytes = 64 << 10 + maxPayloadBytes = 4 << 10 +) + +var ( + ErrInvalidFrame = errors.New("invalid tray protocol frame") + ErrFrameTooLarge = errors.New("tray protocol frame is too large") + ErrUnauthorized = errors.New("tray peer is not authorized for this action") + ErrInvalidPeer = errors.New("tray peer identity is not verified") +) + +type Action uint16 + +const ( + ActionStatus Action = iota + 1 + ActionOpenConfig + ActionOpenLog + ActionStartService + ActionStopService + ActionRestartService + ActionExitTray +) + +func (action Action) valid() bool { return action >= ActionStatus && action <= ActionExitTray } + +// Frame is deliberately not an RPC envelope. Payloads are bounded display +// text only (status/detail); service mutations use an enum and are rechecked +// by the service under the caller's token. +type Frame struct { + Action Action + Payload []byte +} + +func Encode(frame Frame) ([]byte, error) { + if !frame.Action.valid() || len(frame.Payload) > maxPayloadBytes || !utf8.Valid(frame.Payload) { + return nil, ErrInvalidFrame + } + if frame.Action != ActionStatus && len(frame.Payload) != 0 { + return nil, ErrInvalidFrame + } + total := 4 + 2 + 2 + 4 + len(frame.Payload) + if total > maxFrameBytes { + return nil, ErrFrameTooLarge + } + encoded := make([]byte, total) + copy(encoded[:4], []byte("RVTY")) + binary.BigEndian.PutUint16(encoded[4:6], protocolVersion) + binary.BigEndian.PutUint16(encoded[6:8], uint16(frame.Action)) + binary.BigEndian.PutUint32(encoded[8:12], uint32(len(frame.Payload))) + copy(encoded[12:], frame.Payload) + return encoded, nil +} + +func Decode(encoded []byte) (Frame, error) { + if len(encoded) > maxFrameBytes { + return Frame{}, ErrFrameTooLarge + } + if len(encoded) < 12 || !bytes.Equal(encoded[:4], []byte("RVTY")) || binary.BigEndian.Uint16(encoded[4:6]) != protocolVersion { + return Frame{}, ErrInvalidFrame + } + action := Action(binary.BigEndian.Uint16(encoded[6:8])) + length := binary.BigEndian.Uint32(encoded[8:12]) + if !action.valid() || length > maxPayloadBytes || uint64(length)+12 != uint64(len(encoded)) { + return Frame{}, ErrInvalidFrame + } + payload := bytes.Clone(encoded[12:]) + if !utf8.Valid(payload) || action != ActionStatus && len(payload) != 0 { + return Frame{}, ErrInvalidFrame + } + return Frame{Action: action, Payload: payload}, nil +} + +type Peer struct { + PID uint32 + SessionID uint32 + SID string + TokenVerified bool + Interactive bool + Administrator bool + System bool +} + +func (peer Peer) Validate() error { + if peer.PID == 0 || peer.SessionID == ^uint32(0) || peer.SID == "" || !peer.TokenVerified { + return ErrInvalidPeer + } + return nil +} + +func Authorize(peer Peer, action Action) error { + if !action.valid() { + return ErrInvalidFrame + } + if err := peer.Validate(); err != nil { + return err + } + if !peer.Interactive { + return fmt.Errorf("%w: tray peer is not interactive", ErrUnauthorized) + } + switch action { + case ActionStatus, ActionOpenConfig, ActionOpenLog, ActionExitTray: + return nil + case ActionStartService, ActionStopService, ActionRestartService: + if peer.Administrator || peer.System { + return nil + } + return fmt.Errorf("%w: service mutation requires administrator authorization", ErrUnauthorized) + default: + return ErrInvalidFrame + } +} diff --git a/internal/client/windowstray/protocol_test.go b/internal/client/windowstray/protocol_test.go new file mode 100644 index 0000000..5f1319d --- /dev/null +++ b/internal/client/windowstray/protocol_test.go @@ -0,0 +1,56 @@ +package windowstray + +import ( + "bytes" + "errors" + "testing" +) + +func TestTrayFrameRoundTripAndPayloadBounds_HP_WINTRAY_01(t *testing.T) { + t.Parallel() + frame := Frame{Action: ActionStatus, Payload: []byte("connected=true dirty=false")} + encoded, err := Encode(frame) + if err != nil { + t.Fatal(err) + } + decoded, err := Decode(encoded) + if err != nil || decoded.Action != frame.Action || !bytes.Equal(decoded.Payload, frame.Payload) { + t.Fatalf("tray frame round trip = %#v, %v", decoded, err) + } + for _, invalid := range []Frame{{Action: 0}, {Action: ActionOpenLog, Payload: []byte("unexpected")}, {Action: ActionStatus, Payload: bytes.Repeat([]byte("x"), maxPayloadBytes+1)}} { + if !errors.Is(mustEncode(invalid), ErrInvalidFrame) && !errors.Is(mustEncode(invalid), ErrFrameTooLarge) { + t.Fatalf("invalid tray frame accepted: %#v", invalid) + } + } + corrupt := append([]byte(nil), encoded...) + corrupt[0] = 'X' + if _, err := Decode(corrupt); !errors.Is(err, ErrInvalidFrame) { + t.Fatalf("corrupt tray frame error = %v", err) + } +} + +func mustEncode(frame Frame) error { + _, err := Encode(frame) + return err +} + +func TestTrayPeerAuthorizationIsActionScoped_BH_WINTRAY_01(t *testing.T) { + t.Parallel() + user := Peer{PID: 10, SessionID: 1, SID: "S-1-5-21-user", TokenVerified: true, Interactive: true} + if err := Authorize(user, ActionStatus); err != nil { + t.Fatal(err) + } + if err := Authorize(user, ActionRestartService); !errors.Is(err, ErrUnauthorized) { + t.Fatalf("unprivileged service mutation error = %v", err) + } + admin := user + admin.Administrator = true + if err := Authorize(admin, ActionRestartService); err != nil { + t.Fatal(err) + } + for _, peer := range []Peer{{PID: 10, SessionID: 1, SID: "S-1-5-21-user", Interactive: true}, {PID: 10, SessionID: 1, SID: "S-1-5-21-user", TokenVerified: true}} { + if err := Authorize(peer, ActionStatus); !errors.Is(err, ErrInvalidPeer) && !errors.Is(err, ErrUnauthorized) { + t.Fatalf("invalid peer accepted: %#v err=%v", peer, err) + } + } +} diff --git a/internal/domain/reconcile.go b/internal/domain/reconcile.go index 51e78a3..1b81822 100644 --- a/internal/domain/reconcile.go +++ b/internal/domain/reconcile.go @@ -71,9 +71,12 @@ func DecideReconciliation(input ReconcileInput) (ReconcileDecision, error) { if input.ClientEvidence == ClientEvidenceAbsent { return ReconcileDecision{Action: ReconcileNoop}, nil } - if input.ClientEvidence == ClientEvidenceTombstone || IsTerminal(input.ClientLifecycle) || input.ServerHasTombstone { + if input.ClientEvidence == ClientEvidenceTombstone || IsTerminal(input.ClientLifecycle) { return ReconcileDecision{Action: ReconcileDiscardLocalTerminal, EffectiveRevision: input.ClientRevision, SuppressReplay: true}, nil } + if input.ServerHasTombstone { + return ReconcileDecision{Action: ReconcileTerminateLocal, EffectiveRevision: input.ClientRevision, RecordIncident: true, SuppressReplay: true}, nil + } return ReconcileDecision{Action: ReconcileTerminateLocal, EffectiveRevision: input.ClientRevision, RecordIncident: true, SuppressReplay: true}, nil } diff --git a/internal/domain/reconcile_test.go b/internal/domain/reconcile_test.go index af96605..d53c051 100644 --- a/internal/domain/reconcile_test.go +++ b/internal/domain/reconcile_test.go @@ -28,6 +28,7 @@ func TestReconciliationMatrix_HP_SES_11(t *testing.T) { {"terminal active", retainedInput(rvboxv1.CommandLifecycle_COMMAND_SUCCEEDED, rvboxv1.CommandLifecycle_COMMAND_RUNNING), ReconcileTerminateLocal, true}, {"missing active", ReconcileInput{ClientEvidence: ClientEvidenceRetained, ClientLifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, ClientRevision: 1, ImmutableHashMatches: true}, ReconcileTerminateLocal, true}, {"missing terminal", ReconcileInput{ClientEvidence: ClientEvidenceRetained, ClientLifecycle: rvboxv1.CommandLifecycle_COMMAND_FAILED, ClientRevision: 1, ImmutableHashMatches: true}, ReconcileDiscardLocalTerminal, false}, + {"server tombstone active client", ReconcileInput{ClientEvidence: ClientEvidenceRetained, ClientLifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, ClientRevision: 1, ImmutableHashMatches: true, ServerHasTombstone: true}, ReconcileTerminateLocal, true}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { diff --git a/internal/server/control/service_test.go b/internal/server/control/service_test.go index 1396a4c..3dcef01 100644 --- a/internal/server/control/service_test.go +++ b/internal/server/control/service_test.go @@ -289,7 +289,7 @@ func TestStorageIncidentControlLifecycle_HP_CONTROL_14(t *testing.T) { service, persistence := newTestService(t) defer persistence.Close() incidentID := fixedIssue(0xac) - if _, err := persistence.RecordIncident(context.Background(), store.IncidentInput{IncidentUUID: [16]byte(incidentID), DetectedAt: time.Now().UTC(), Kind: store.IncidentChecksumMismatch, Scope: store.IncidentScopeGlobal, ScopeKey: "segments", Summary: "checksum mismatch", Evidence: []byte("evidence"), AutomaticallyRepairable: true}); err != nil { + if _, err := persistence.RecordIncident(context.Background(), store.IncidentInput{IncidentUUID: [16]byte(incidentID), DetectedAt: time.Date(2026, time.September, 6, 11, 59, 0, 0, time.UTC), Kind: store.IncidentChecksumMismatch, Scope: store.IncidentScopeGlobal, ScopeKey: "segments", Summary: "checksum mismatch", Evidence: []byte("evidence"), AutomaticallyRepairable: true}); err != nil { t.Fatal(err) } listed, err := service.ListStorageIncidents(context.Background(), &rvboxv1.ListStorageIncidentsRequest{PageSize: 1}) diff --git a/internal/server/session/agent_server.go b/internal/server/session/agent_server.go index 9e49787..83a629f 100644 --- a/internal/server/session/agent_server.go +++ b/internal/server/session/agent_server.go @@ -180,6 +180,19 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w server.close(connection, websocket.StatusInternalError, "could not queue welcome") return } + targets, err := server.Store.ReconcileTargets(parent, hello.GetClientId()) + if err != nil { + server.close(connection, websocket.StatusInternalError, "could not build reconciliation request") + return + } + reconcileRequest, err := proto.Marshal(&rvboxv1.AgentEnvelope{ + SessionId: encodedSessionID, SessionGeneration: registration.Generation, + Payload: &rvboxv1.AgentEnvelope_ReconcileRequest{ReconcileRequest: &rvboxv1.ReconcileRequest{Targets: targets}}, + }) + if err != nil || queue.EnqueueControl(Frame{Kind: FrameControl, Payload: reconcileRequest}) != nil { + server.close(connection, websocket.StatusInternalError, "could not queue reconciliation request") + return + } heartbeat := newSynchronizedHeartbeat(server.heartbeatIdle(), server.livenessTimeout(), 0) writerDone := make(chan struct{}) @@ -213,7 +226,7 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w return } if snapshot := envelope.GetReconcileSnapshot(); snapshot != nil { - result, reconcileErr := server.Store.ReconcileClientSnapshot(sessionContext, hello.GetClientId(), snapshot) + result, reconcileErr := server.Store.ReconcileClientSnapshotForSession(sessionContext, hello.GetClientId(), registration.Generation, snapshot) if reconcileErr != nil { server.close(connection, websocket.StatusPolicyViolation, "reconciliation failed") return diff --git a/internal/server/session/agent_server_test.go b/internal/server/session/agent_server_test.go index 0cce729..988c1a9 100644 --- a/internal/server/session/agent_server_test.go +++ b/internal/server/session/agent_server_test.go @@ -238,6 +238,25 @@ func dialAndHello(t *testing.T, server *httptest.Server, clientID, instanceID st connection.CloseNow() t.Fatalf("welcome payload = %T", welcome.Payload) } + // Reconnecting agents receive an advisory target list immediately after + // welcome. Consume it here so callers that exercise fencing or malformed + // frames observe the next read from the actual session boundary. + readContext, cancel = context.WithTimeout(context.Background(), time.Second) + defer cancel() + messageType, advisoryBytes, err := connection.Read(readContext) + if err != nil { + connection.CloseNow() + t.Fatal(err) + } + if messageType != websocket.MessageBinary { + connection.CloseNow() + t.Fatalf("reconciliation message type = %v", messageType) + } + var advisory rvboxv1.AgentEnvelope + if err := proto.Unmarshal(advisoryBytes, &advisory); err != nil || advisory.GetReconcileRequest() == nil { + connection.CloseNow() + t.Fatalf("reconciliation advisory = %+v, %v", advisory.Payload, err) + } return connection, &welcome } diff --git a/internal/server/store/reconcile.go b/internal/server/store/reconcile.go index 7afdf89..82d708d 100644 --- a/internal/server/store/reconcile.go +++ b/internal/server/store/reconcile.go @@ -4,60 +4,302 @@ import ( "context" "database/sql" "errors" + "fmt" + "sort" + "time" rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/domain" ) +// ReconcileTargets returns the server's complete non-terminal view for the +// client. It is a read-only snapshot used to tell a reconnecting agent which +// UUIDs and durable cursors must be compared before fresh dispatch is enabled. +func (store *Store) ReconcileTargets(ctx context.Context, clientID string) ([]*rvboxv1.ReconcileTarget, error) { + if clientID == "" { + return nil, errors.New("client ID is required") + } + rows, err := store.db.QueryContext(ctx, `SELECT issue_uuid, last_event_seq, revision, immutable_request_sha256 +FROM commands WHERE client_id = ? AND lifecycle BETWEEN 1 AND 4 ORDER BY issue_time, issue_uuid`, clientID) + if err != nil { + return nil, err + } + defer rows.Close() + var result []*rvboxv1.ReconcileTarget + for rows.Next() { + var issue, digest []byte + var eventSeq, revision uint64 + if err := rows.Scan(&issue, &eventSeq, &revision, &digest); err != nil { + return nil, err + } + if len(issue) != 16 || len(digest) != 32 || revision == 0 { + return nil, ErrInvalidSegmentRecord + } + var parsed domain.UUID + copy(parsed[:], issue) + if _, err := domain.ParseUUIDv7(parsed.String()); err != nil { + return nil, ErrInvalidSegmentRecord + } + result = append(result, &rvboxv1.ReconcileTarget{IssueUuid: parsed.String(), LastServerEventSeq: eventSeq, CommandRevision: revision, ImmutableRequestSha256: append([]byte(nil), digest...)}) + } + if err := rows.Err(); err != nil { + return nil, err + } + return result, nil +} + // ReconcileClientSnapshot compares client evidence with this client's durable // server rows. It intentionally performs no lifecycle mutation yet: callers // receive only the safe local terminate/discard instructions. func (store *Store) ReconcileClientSnapshot(ctx context.Context, clientID string, snapshot *rvboxv1.ReconcileSnapshot) (*rvboxv1.ReconcileResult, error) { + return store.reconcileClientSnapshot(ctx, clientID, 0, snapshot) +} + +// ReconcileClientSnapshotForSession applies the complete bidirectional +// reconciliation matrix while the registering session is fenced. Queued or +// dispatched rows absent from the client's complete snapshot are safely +// requeued; accepted/running rows absent from the snapshot are interrupted and +// incidented. Retained non-terminal rows are retargeted to this generation so +// late events from the previous connection cannot advance the command. +func (store *Store) ReconcileClientSnapshotForSession(ctx context.Context, clientID string, generation uint64, snapshot *rvboxv1.ReconcileSnapshot) (*rvboxv1.ReconcileResult, error) { + if generation == 0 { + return nil, errors.New("session generation is required") + } + return store.reconcileClientSnapshot(ctx, clientID, generation, snapshot) +} + +type reconcileServerRow struct { + issue domain.UUID + lifecycle rvboxv1.CommandLifecycle + revision uint64 + lastSeq uint64 + hash []byte + target sql.NullInt64 +} + +func (store *Store) reconcileClientSnapshot(ctx context.Context, clientID string, generation uint64, snapshot *rvboxv1.ReconcileSnapshot) (*rvboxv1.ReconcileResult, error) { + if clientID == "" { + return nil, errors.New("client ID is required") + } if snapshot == nil { return nil, errors.New("missing client reconciliation snapshot") } - result := &rvboxv1.ReconcileResult{} + clientRows := make(map[string]*rvboxv1.ReconcileCommandState, len(snapshot.GetRetainedCommands())) for _, client := range snapshot.GetRetainedCommands() { + if client == nil || client.GetCommandRevision() == 0 || len(client.GetImmutableRequestSha256()) != 32 { + return nil, fmt.Errorf("invalid client reconciliation row") + } issue, err := domain.ParseUUIDv7(client.GetIssueUuid()) if err != nil { return nil, err } - input := domain.ReconcileInput{ClientEvidence: domain.ClientEvidenceRetained, ClientLifecycle: client.GetLifecycle(), ClientRevision: domain.CommandRevision(client.GetCommandRevision()), ClientLastEventSeq: client.GetLastClientEventSeq(), ImmutableHashMatches: true} + key := issue.String() + if _, exists := clientRows[key]; exists { + return nil, fmt.Errorf("duplicate client reconciliation row %s", key) + } + clientRows[key] = client + } + + store.writeMu.Lock() + database, err := store.openDatabase() + if err != nil { + store.writeMu.Unlock() + return nil, err + } + tx, err := database.BeginTx(ctx, nil) + if err != nil { + store.writeMu.Unlock() + return nil, err + } + defer tx.Rollback() + serverRows, tombstones, err := loadReconcileRows(ctx, tx, clientID) + if err != nil { + store.writeMu.Unlock() + return nil, err + } + result := &rvboxv1.ReconcileResult{} + var incidents []reconcileIncident + for key, client := range clientRows { + row, present := serverRows[key] + input := domain.ReconcileInput{ClientEvidence: domain.ClientEvidenceRetained, ClientLifecycle: client.GetLifecycle(), ClientRevision: domain.CommandRevision(client.GetCommandRevision()), ClientLastEventSeq: client.GetLastClientEventSeq(), ImmutableHashMatches: false} if client.GetTombstoned() { input.ClientEvidence = domain.ClientEvidenceTombstone } - var lifecycle uint32 - var revision, lastSequence uint64 - var hash []byte - err = store.db.QueryRowContext(ctx, `SELECT lifecycle, revision, last_event_seq, immutable_request_sha256 FROM commands WHERE issue_uuid = ? AND client_id = ?`, issue[:], clientID).Scan(&lifecycle, &revision, &lastSequence, &hash) - if err == nil { + if present { input.ServerPresent = true - input.ServerLifecycle = rvboxv1.CommandLifecycle(lifecycle) - input.ServerRevision = domain.CommandRevision(revision) - input.ServerLastEventSeq = lastSequence - input.ImmutableHashMatches = len(hash) == len(client.GetImmutableRequestSha256()) && string(hash) == string(client.GetImmutableRequestSha256()) - } else if !errors.Is(err, sql.ErrNoRows) { - return nil, err + input.ServerLifecycle = row.lifecycle + input.ServerRevision = domain.CommandRevision(row.revision) + input.ServerLastEventSeq = row.lastSeq + input.ImmutableHashMatches = len(row.hash) == 32 && bytesEqual(row.hash, client.GetImmutableRequestSha256()) + } else if tombstoneHash, found := tombstones[key]; found { + input.ServerHasTombstone = true + input.ImmutableHashMatches = len(tombstoneHash) == 32 && bytesEqual(tombstoneHash, client.GetImmutableRequestSha256()) } else { - var tombstoneHash []byte - err = store.db.QueryRowContext(ctx, `SELECT immutable_sha256 FROM command_tombstones WHERE issue_uuid = ? AND client_id = ?`, issue[:], clientID).Scan(&tombstoneHash) - if err == nil { - input.ServerHasTombstone = true - input.ImmutableHashMatches = len(tombstoneHash) == len(client.GetImmutableRequestSha256()) && string(tombstoneHash) == string(client.GetImmutableRequestSha256()) - } else if !errors.Is(err, sql.ErrNoRows) { - return nil, err - } + // Once the server has neither a live row nor a tombstone, there is + // no immutable request hash left to compare. The client evidence is + // still useful: a retained non-terminal row must be terminated and + // a retained terminal row may be discarded. Treat the evidence as + // structurally valid here; hash equality is required whenever the + // server still has a row or tombstone to compare against. + input.ImmutableHashMatches = true } - decision, err := domain.DecideReconciliation(input) - if err != nil && !errors.Is(err, domain.ErrReconcileContradiction) { - return nil, err + decision, decisionErr := domain.DecideReconciliation(input) + if decisionErr != nil { + incidents = append(incidents, reconcileIncident{issue: parseIssueUnchecked(key), summary: decisionErr.Error(), dataLoss: false}) + continue } switch decision.Action { case domain.ReconcileTerminateLocal: - result.TerminateLocalIssueUuids = append(result.TerminateLocalIssueUuids, issue.String()) + result.TerminateLocalIssueUuids = append(result.TerminateLocalIssueUuids, key) + if decision.RecordIncident { + incidents = append(incidents, reconcileIncident{issue: parseIssueUnchecked(key), summary: "client and server command state disagree; local process must be terminated", dataLoss: true}) + } case domain.ReconcileDiscardLocalTerminal: - result.DiscardLocalTerminalIssueUuids = append(result.DiscardLocalTerminalIssueUuids, issue.String()) + result.DiscardLocalTerminalIssueUuids = append(result.DiscardLocalTerminalIssueUuids, key) + case domain.ReconcileResumeDelivery: + if generation != 0 && row.target.Valid && uint64(row.target.Int64) != generation || generation != 0 && !row.target.Valid { + if _, err := tx.ExecContext(ctx, `UPDATE commands SET target_session_generation = ? WHERE issue_uuid = ? AND client_id = ? AND lifecycle BETWEEN 2 AND 4`, generation, row.issue[:], clientID); err != nil { + store.writeMu.Unlock() + return nil, err + } + } + case domain.ReconcileInterruptServerSuppressReplay: + if err := interruptReconcileRow(ctx, tx, row, clientID, "client tombstone contradicts non-terminal server command"); err != nil { + store.writeMu.Unlock() + return nil, err + } + incidents = append(incidents, reconcileIncident{issue: row.issue, summary: "client tombstone contradicts non-terminal server command", dataLoss: true}) } } + for key, row := range serverRows { + if _, present := clientRows[key]; present { + continue + } + decision, decisionErr := domain.DecideReconciliation(domain.ReconcileInput{ServerPresent: true, ServerLifecycle: row.lifecycle, ServerRevision: domain.CommandRevision(row.revision), ServerLastEventSeq: row.lastSeq, ClientEvidence: domain.ClientEvidenceAbsent}) + if decisionErr != nil { + incidents = append(incidents, reconcileIncident{issue: row.issue, summary: decisionErr.Error(), dataLoss: false}) + continue + } + switch decision.Action { + case domain.ReconcileRequeue: + if _, err := tx.ExecContext(ctx, `UPDATE commands SET lifecycle = 1, target_session_generation = NULL WHERE issue_uuid = ? AND client_id = ? AND lifecycle IN (1, 2)`, row.issue[:], clientID); err != nil { + store.writeMu.Unlock() + return nil, err + } + case domain.ReconcileInterruptServerClientStateLoss: + if err := interruptReconcileRow(ctx, tx, row, clientID, "client lost accepted/running command state"); err != nil { + store.writeMu.Unlock() + return nil, err + } + incidents = append(incidents, reconcileIncident{issue: row.issue, summary: "client lost accepted/running command state", dataLoss: true}) + } + } + if err := tx.Commit(); err != nil { + store.writeMu.Unlock() + return nil, err + } + store.writeMu.Unlock() + for _, incident := range incidents { + if err := store.recordReconcileIncident(ctx, clientID, incident); err != nil { + return nil, err + } + } + sort.Strings(result.TerminateLocalIssueUuids) + sort.Strings(result.DiscardLocalTerminalIssueUuids) return result, nil } + +type reconcileIncident struct { + issue domain.UUID + summary string + dataLoss bool +} + +func loadReconcileRows(ctx context.Context, tx *sql.Tx, clientID string) (map[string]reconcileServerRow, map[string][]byte, error) { + rows, err := tx.QueryContext(ctx, `SELECT issue_uuid, lifecycle, revision, last_event_seq, immutable_request_sha256, target_session_generation FROM commands WHERE client_id = ?`, clientID) + if err != nil { + return nil, nil, err + } + commands := make(map[string]reconcileServerRow) + for rows.Next() { + var encoded, hash []byte + var lifecycle uint32 + var row reconcileServerRow + if err := rows.Scan(&encoded, &lifecycle, &row.revision, &row.lastSeq, &hash, &row.target); err != nil { + _ = rows.Close() + return nil, nil, err + } + if len(encoded) != 16 || len(hash) != 32 || lifecycle == 0 || lifecycle > 11 || row.revision == 0 { + _ = rows.Close() + return nil, nil, ErrInvalidSegmentRecord + } + copy(row.issue[:], encoded) + if _, err := domain.ParseUUIDv7(row.issue.String()); err != nil { + return nil, nil, ErrInvalidSegmentRecord + } + row.lifecycle = rvboxv1.CommandLifecycle(lifecycle) + row.hash = append([]byte(nil), hash...) + commands[row.issue.String()] = row + } + if err := rows.Err(); err != nil { + _ = rows.Close() + return nil, nil, err + } + if err := rows.Close(); err != nil { + return nil, nil, err + } + tombstoneRows, err := tx.QueryContext(ctx, `SELECT issue_uuid, immutable_sha256 FROM command_tombstones WHERE client_id = ?`, clientID) + if err != nil { + return nil, nil, err + } + defer tombstoneRows.Close() + tombstones := make(map[string][]byte) + for tombstoneRows.Next() { + var encoded, hash []byte + if err := tombstoneRows.Scan(&encoded, &hash); err != nil { + return nil, nil, err + } + var issue domain.UUID + if len(encoded) != 16 || len(hash) != 32 { + return nil, nil, ErrInvalidSegmentRecord + } + copy(issue[:], encoded) + tombstones[issue.String()] = append([]byte(nil), hash...) + } + return commands, tombstones, tombstoneRows.Err() +} + +func interruptReconcileRow(ctx context.Context, tx *sql.Tx, row reconcileServerRow, clientID, detail string) error { + _, err := tx.ExecContext(ctx, `UPDATE commands SET lifecycle = 9, terminal_time = ?, revision = revision + 1, target_session_generation = NULL WHERE issue_uuid = ? AND client_id = ? AND lifecycle BETWEEN 2 AND 4`, time.Now().UTC().UnixNano(), row.issue[:], clientID) + return err +} + +func parseIssueUnchecked(value string) domain.UUID { + issue, _ := domain.ParseUUIDv7(value) + return issue +} + +func bytesEqual(left, right []byte) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] != right[index] { + return false + } + } + return true +} + +func (store *Store) recordReconcileIncident(ctx context.Context, clientID string, incident reconcileIncident) error { + if incident.issue == (domain.UUID{}) { + return errors.New("invalid reconciliation incident issue") + } + id, err := domain.NewUUIDv7() + if err != nil { + return err + } + issue := [16]byte(incident.issue) + _, err = store.RecordIncident(ctx, IncidentInput{IncidentUUID: [16]byte(id), DetectedAt: time.Now().UTC(), Kind: IncidentCounterMismatch, Scope: IncidentScopeCommand, ScopeKey: incident.issue.String(), ClientID: clientID, IssueUUID: &issue, Summary: incident.summary, Evidence: []byte(incident.summary), DataLoss: incident.dataLoss, AutomaticallyRepairable: false}) + return err +} diff --git a/internal/server/store/reconcile_test.go b/internal/server/store/reconcile_test.go new file mode 100644 index 0000000..6712ddb --- /dev/null +++ b/internal/server/store/reconcile_test.go @@ -0,0 +1,81 @@ +package store + +import ( + "context" + "crypto/sha256" + "path/filepath" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/domain" +) + +func TestReconcileSnapshotMutatesMissingAndRetargets_HP_SES_13(t *testing.T) { + ctx := context.Background() + opened, err := Open(ctx, Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = opened.Close() }) + if _, err := opened.RegisterClientSession(ctx, ClientRegistration{ClientID: "reconcile-client", Platform: 3, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\`, SupportedShells: []byte{1}, ClientInstanceID: [16]byte{1}, SessionID: [16]byte{2}, ConnectedAt: time.Now().UTC()}); err != nil { + t.Fatal(err) + } + if _, err := opened.RegisterClientSession(ctx, ClientRegistration{ClientID: "reconcile-active", Platform: 3, Architecture: "amd64", DaemonVersion: "test", DaemonCWD: `C:\`, SupportedShells: []byte{1}, ClientInstanceID: [16]byte{3}, SessionID: [16]byte{4}, ConnectedAt: time.Now().UTC()}); err != nil { + t.Fatal(err) + } + now := time.Now().UTC() + queued := mustReconcileIssue(t, "019c46f1-1d02-7000-8000-0000000000c1") + if _, err := opened.QueueCommand(ctx, QueueCommandInput{IssueUUID: queued, ClientID: "reconcile-client", IssueTime: now, ReceiptTime: now, ImmutableSHA256: sha256.Sum256([]byte("queued")), ExecutionSpec: []byte("opaque")}); err != nil { + t.Fatal(err) + } + if _, err := opened.ClaimNextDispatch(ctx, "reconcile-client", 1, now); err != nil { + t.Fatal(err) + } + result, err := opened.ReconcileClientSnapshotForSession(ctx, "reconcile-client", 2, &rvboxv1.ReconcileSnapshot{}) + if err != nil || len(result.GetTerminateLocalIssueUuids()) != 0 { + t.Fatalf("missing dispatched result = %#v, %v", result, err) + } + var lifecycle uint32 + var target any + if err := opened.DB().QueryRowContext(ctx, `SELECT lifecycle, target_session_generation FROM commands WHERE issue_uuid = ?`, queued[:]).Scan(&lifecycle, &target); err != nil { + t.Fatal(err) + } + if lifecycle != uint32(rvboxv1.CommandLifecycle_COMMAND_QUEUED) || target != nil { + t.Fatalf("missing dispatched command = lifecycle=%d target=%v", lifecycle, target) + } + + active := mustReconcileIssue(t, "019c46f1-1d02-7000-8000-0000000000c2") + if _, err := opened.QueueCommand(ctx, QueueCommandInput{IssueUUID: active, ClientID: "reconcile-active", IssueTime: now.Add(time.Second), ReceiptTime: now.Add(time.Second), ImmutableSHA256: sha256.Sum256([]byte("active")), ExecutionSpec: []byte("opaque")}); err != nil { + t.Fatal(err) + } + if _, err := opened.ClaimNextDispatch(ctx, "reconcile-active", 1, now.Add(time.Second)); err != nil { + t.Fatal(err) + } + if _, err := opened.RecordCommandAcceptance(ctx, active, "reconcile-active", 1, 1, true, now.Add(2*time.Second)); err != nil { + t.Fatal(err) + } + result, err = opened.ReconcileClientSnapshotForSession(ctx, "reconcile-active", 2, &rvboxv1.ReconcileSnapshot{}) + if err != nil || len(result.GetTerminateLocalIssueUuids()) != 0 { + t.Fatalf("missing accepted result = %#v, %v", result, err) + } + if err := opened.DB().QueryRowContext(ctx, `SELECT lifecycle, revision, target_session_generation FROM commands WHERE issue_uuid = ?`, active[:]).Scan(&lifecycle, new(uint64), &target); err != nil { + t.Fatal(err) + } + if lifecycle != uint32(rvboxv1.CommandLifecycle_COMMAND_INTERRUPTED) || target != nil { + t.Fatalf("missing accepted command = lifecycle=%d target=%v", lifecycle, target) + } + var incidents int + if err := opened.DB().QueryRowContext(ctx, `SELECT count(*) FROM storage_incidents WHERE scope = 'command' AND scope_key = ?`, active.String()).Scan(&incidents); err != nil || incidents != 1 { + t.Fatalf("reconciliation incident count = %d, %v", incidents, err) + } +} + +func mustReconcileIssue(t *testing.T, value string) domain.UUID { + t.Helper() + issue, err := domain.ParseUUIDv7(value) + if err != nil { + t.Fatal(err) + } + return issue +} diff --git a/scripts/test-e2e b/scripts/test-e2e index 78dfa67..929e42e 100755 --- a/scripts/test-e2e +++ b/scripts/test-e2e @@ -1,6 +1,12 @@ #!/bin/sh set -eu -echo "E2E requires the production server plus a native Windows host; that host is currently unavailable" >&2 -exit 2 +repo_root=$(CDPATH= cd -- "$(dirname -- "$0")/.." && pwd) +# The deterministic protocol E2E lane runs in the pinned toolchain container. +# A native Windows run is opt-in and is driven by the checked-in PowerShell +# host adapter; it must acquire its VM lease and reset the exact fixture before +# mutating it. No guest password is accepted on this command line. +cd "$repo_root" +exec docker compose -f deploy/compose.yaml run --rm toolchain \ + go run ./test/harness e2e "$@" diff --git a/scripts/windows/test-host.ps1 b/scripts/windows/test-host.ps1 new file mode 100644 index 0000000..10c9729 --- /dev/null +++ b/scripts/windows/test-host.ps1 @@ -0,0 +1,103 @@ +[CmdletBinding()] +param( + [Parameter(Mandatory = $true, Position = 0)] + [ValidateSet('Prepare', 'Status', 'Run', 'Collect', 'Stop', 'Reset')] + [string] $Action, + [Parameter(Mandatory = $false)] + [ValidatePattern('^[a-z0-9][a-z0-9-]{0,63}$')] + [string] $RunId = '', + [string] $VmName = $(if ($env:RVBOX_WINDOWS_VM) { $env:RVBOX_WINDOWS_VM } else { 'rvbox-win10-test' }), + [string] $Snapshot = $(if ($env:RVBOX_WINDOWS_BASELINE_SNAPSHOT) { $env:RVBOX_WINDOWS_BASELINE_SNAPSHOT } else { 'baseline-disk-first' }) +) + +$ErrorActionPreference = 'Stop' +$VBoxManage = if ($env:VBOXMANAGE) { $env:VBOXMANAGE } else { 'VBoxManage' } +$GuestUser = $env:RVBOX_WINDOWS_GUEST_USER +$GuestPassword = $env:RVBOX_WINDOWS_GUEST_PASSWORD +$LeaseRoot = if ($env:RVBOX_WINDOWS_LEASE_DIR) { $env:RVBOX_WINDOWS_LEASE_DIR } else { Join-Path $PSScriptRoot '..\..\.test-runs\windows' } + +function Invoke-VBox([string[]] $Arguments) { + $output = & $VBoxManage @Arguments 2>&1 + if ($LASTEXITCODE -ne 0) { + throw "VBoxManage failed ($LASTEXITCODE): $($output -join ' ')" + } + return $output +} + +function Require-RunId { + if ([string]::IsNullOrWhiteSpace($RunId)) { throw "$Action requires -RunId" } +} + +function Acquire-Lease { + New-Item -ItemType Directory -Force -Path $LeaseRoot | Out-Null + $path = Join-Path $LeaseRoot 'lease.lock' + try { + $script:Lease = [System.IO.File]::Open($path, [System.IO.FileMode]::OpenOrCreate, [System.IO.FileAccess]::ReadWrite, [System.IO.FileShare]::None) + $bytes = [Text.Encoding]::UTF8.GetBytes("$VmName`n$RunId`n$([DateTime]::UtcNow.ToString('o'))`n") + $Lease.SetLength(0); $Lease.Write($bytes, 0, $bytes.Length); $Lease.Flush($true) + } catch { + throw "Windows test VM is already leased: $path" + } +} + +function Release-Lease { + if ($script:Lease) { $script:Lease.Dispose(); $script:Lease = $null } +} + +function Guest-Args([string[]] $Arguments) { + if ([string]::IsNullOrWhiteSpace($GuestUser) -or [string]::IsNullOrWhiteSpace($GuestPassword)) { + throw 'Set RVBOX_WINDOWS_GUEST_USER and RVBOX_WINDOWS_GUEST_PASSWORD in the host environment; secrets are never read from repository files or printed.' + } + return @('guestcontrol', $VmName, '--username', $GuestUser, '--password', $GuestPassword) + $Arguments +} + +try { + if ($Action -eq 'Status') { + Invoke-VBox @('showvminfo', $VmName, '--machinereadable') | Where-Object { $_ -match '^(VMState|name|SnapshotName)=' } + exit 0 + } + Require-RunId + Acquire-Lease + try { + $state = (Invoke-VBox @('showvminfo', $VmName, '--machinereadable') | Where-Object { $_ -like 'VMState=*' }) -replace '^VMState="?([^"\r\n]+)"?$', '$1' + switch ($Action) { + 'Prepare' { + if ($state -eq 'running') { Invoke-VBox (Guest-Args @('run', 'C:\Windows\System32\whoami.exe', '--', '/groups')) | Out-Null } + else { + Invoke-VBox @('startvm', $VmName, '--type', 'headless') | Out-Null + Start-Sleep -Seconds 5 + Invoke-VBox (Guest-Args @('run', 'C:\Windows\System32\whoami.exe', '--', '/groups')) | Out-Null + } + Invoke-VBox (Guest-Args @('run', 'C:\Windows\System32\query.exe', '--', 'user')) | Out-Null + Write-Output "prepared VM=$VmName run=$RunId" + } + 'Run' { + $guestBinary = if ($env:RVBOX_WINDOWS_GUEST_BINARY) { $env:RVBOX_WINDOWS_GUEST_BINARY } else { 'C:\ProgramData\RVBox\test\rvbox.exe' } + Invoke-VBox (Guest-Args @('run', $guestBinary, '--', '--check-config', '--config', 'C:\ProgramData\RVBox\client.toml')) | Out-Null + Write-Output "ran native Windows smoke run=$RunId" + } + 'Collect' { + $destination = Join-Path $LeaseRoot $RunId + New-Item -ItemType Directory -Force -Path $destination | Out-Null + $guestArtifacts = if ($env:RVBOX_WINDOWS_GUEST_ARTIFACTS) { $env:RVBOX_WINDOWS_GUEST_ARTIFACTS } else { 'C:\ProgramData\RVBox\test-artifacts' } + Invoke-VBox (Guest-Args @('copyfrom', $guestArtifacts, $destination, '--recursive')) | Out-Null + Write-Output "collected native artifacts under $destination" + } + 'Stop' { + if ($state -eq 'running') { Invoke-VBox @('controlvm', $VmName, 'acpipowerbutton') | Out-Null } + Write-Output "requested graceful stop VM=$VmName" + } + 'Reset' { + if ($state -eq 'running') { Invoke-VBox @('controlvm', $VmName, 'acpipowerbutton') | Out-Null; Start-Sleep -Seconds 3 } + Invoke-VBox @('snapshot', $VmName, 'restore', $Snapshot) | Out-Null + Write-Output "restored baseline snapshot=$Snapshot VM=$VmName run=$RunId" + } + } + } finally { + Release-Lease + } +} catch { + Release-Lease + Write-Error $_ + exit 1 +} diff --git a/test/coverage.toml b/test/coverage.toml index 8c59988..f65daa2 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -221,6 +221,30 @@ tests = [ "internal/client/supervisor/windows/selection_test.go:TestSelectExecutionContext_HP_WINCTX_02", ] +[[requirements]] +id = "HP-WINSVC-01" +layer = "unit" +status = "implemented" +tests = ["internal/client/windowsservice/service_test.go:TestInstallSpecValidation_HP_WINSVC_01"] + +[[requirements]] +id = "HP-WINSVC-02" +layer = "unit" +status = "implemented" +tests = ["internal/client/windowsservice/service_test.go:TestServiceTransitionsAreIdempotentAndExplicit_HP_WINSVC_02"] + +[[requirements]] +id = "HP-WINTRAY-01" +layer = "unit" +status = "implemented" +tests = ["internal/client/windowstray/protocol_test.go:TestTrayFrameRoundTripAndPayloadBounds_HP_WINTRAY_01"] + +[[requirements]] +id = "BH-WINTRAY-01" +layer = "unit" +status = "implemented" +tests = ["internal/client/windowstray/protocol_test.go:TestTrayPeerAuthorizationIsActionScoped_BH_WINTRAY_01"] + [[requirements]] id = "BH-WINCTX-02" layer = "unit" @@ -245,6 +269,48 @@ layer = "unit" status = "implemented" tests = ["internal/client/supervisor/supervisor_test.go:TestStartSpecValidation_BH_SUPERVISOR_01"] +[[requirements]] +id = "HP-SUPERVISOR-03" +layer = "unit" +status = "implemented" +tests = ["internal/client/supervisor/windows/native_other_test.go:TestPortableSupervisorCapturesOutputAndSupportsStdin_HP_SUPERVISOR_03"] + +[[requirements]] +id = "HP-SUPERVISOR-04" +layer = "unit" +status = "implemented" +tests = ["internal/client/supervisor/windows/native_other_test.go:TestPortableSupervisorScriptMaterializationAndCleanup_HP_SUPERVISOR_04"] + +[[requirements]] +id = "HP-EXECUTOR-01" +layer = "unit" +status = "implemented" +tests = ["internal/client/agent/executor_test.go:TestExecutorRunsDurableCommandAndPublishesTerminalEvents_HP_EXECUTOR_01"] + +[[requirements]] +id = "HP-EXECUTOR-02" +layer = "unit" +status = "implemented" +tests = ["internal/client/agent/executor_test.go:TestExecutorWaitsForScriptCommitBeforeLaunch_HP_EXECUTOR_02"] + +[[requirements]] +id = "HP-RUNTIME-03" +layer = "unit" +status = "implemented" +tests = ["internal/client/agent/runner_test.go:TestFlushEventsAssignsAndSendsOnlyUnacknowledgedRows_HP_RUNTIME_03"] + +[[requirements]] +id = "HP-CLIENT-11" +layer = "unit" +status = "implemented" +tests = ["internal/client/spool/spool_test.go:TestAppendLifecycleAtomicallyUpdatesPhaseAndEvent_HP_CLIENT_11"] + +[[requirements]] +id = "HP-CLIENT-12" +layer = "unit" +status = "implemented" +tests = ["internal/client/spool/recovery_test.go:TestRecoverLaunchUncertaintyFencesRedispatchAfterRestart_HP_CLIENT_12"] + [[requirements]] id = "BH-SES-01" layer = "unit" diff --git a/test/harness/harness.go b/test/harness/harness.go index 78be4c9..ba0873b 100644 --- a/test/harness/harness.go +++ b/test/harness/harness.go @@ -86,6 +86,8 @@ func runCLI(ctx context.Context, args []string) error { return validateCoverageInventory(filepath.Join(repoRoot, "test", "coverage.toml")) case "integration": return h.integration(ctx, args[1:]) + case "e2e": + return h.e2e(ctx, args[1:]) case "status", "logs", "collect", "recover", "reuse", "stop", "reset", "purge": return h.environmentCommand(args[0], args[1:]) default: @@ -94,7 +96,7 @@ func runCLI(ctx context.Context, args []string) error { } func usageError() error { - return errors.New("usage: harness doctor|coverage|integration|status|logs|collect|recover|reuse|stop|reset|purge") + return errors.New("usage: harness doctor|coverage|integration|e2e|status|logs|collect|recover|reuse|stop|reset|purge") } func (h *harness) doctor(repoRoot string) error { @@ -113,7 +115,7 @@ func (h *harness) doctor(repoRoot string) error { name := probe.Name() _ = probe.Close() _ = os.Remove(name) - fmt.Fprintln(h.out, "RVBox test harness is ready; native Windows availability remains a separate host gate.") + fmt.Fprintln(h.out, "RVBox test harness is ready; native Windows scenarios use the explicit scripts/windows/test-host.ps1 host lane.") return nil } @@ -182,6 +184,83 @@ func (h *harness) integration(ctx context.Context, args []string) error { return h.transition(current, "completed", *suite+"-complete", *suite+" integration run completed") } +// e2e runs production-shaped Go scenarios against real local listeners and +// durable stores. Native Windows work is an additional host lane invoked by +// scripts/windows/test-host.ps1; it is never silently replaced by Wine or a +// cross-compiled binary. The run manifest/journal makes every scenario +// resumable and keeps artifacts bounded. +func (h *harness) e2e(ctx context.Context, args []string) error { + flags := flag.NewFlagSet("e2e", flag.ContinueOnError) + flags.SetOutput(io.Discard) + scenario := flags.String("scenario", "smoke", "scenario name: smoke, script, recovery, or all") + runID := flags.String("run-id", "", "run ID") + resume := flags.Bool("resume", false, "resume an existing run") + if err := flags.Parse(args); err != nil { + return err + } + if *scenario != "smoke" && *scenario != "script" && *scenario != "recovery" && *scenario != "all" { + return fmt.Errorf("scenario %q is not implemented yet; available: smoke, script, recovery, all", *scenario) + } + var current *manifest + var err error + if *resume { + if *runID == "" { + return errors.New("--resume requires --run-id") + } + current, err = h.load(*runID) + if err != nil { + return err + } + if current.Layer != "e2e" || current.Suite != *scenario { + return errors.New("run layer/scenario does not match resume request") + } + if current.Phase != "ready" && current.Phase != "interrupted" && current.Phase != "stopped" && current.Phase != "running" && current.Phase != "failed" { + return fmt.Errorf("run in phase %q is not resumable; recover or reuse it first", current.Phase) + } + } else { + current, err = h.create(*runID, "e2e", *scenario) + if err != nil { + return err + } + } + fmt.Fprintln(h.out, current.RunID) + if err := h.transition(current, "running", "e2e-start", "e2e scenario started"); err != nil { + return err + } + scenarios := []string{*scenario} + if *scenario == "all" { + scenarios = []string{"smoke", "script", "recovery"} + } + for _, item := range scenarios { + if err := h.runE2EScenario(ctx, current, item); err != nil { + _ = h.transition(current, "failed", "e2e-"+item+"-failed", err.Error()) + return err + } + } + return h.transition(current, "completed", "e2e-complete", "e2e scenario completed") +} + +func (h *harness) runE2EScenario(ctx context.Context, current *manifest, scenario string) error { + if err := h.appendJournal(current.RunID, journalEntry{At: h.now(), Step: "e2e-" + scenario, Status: "running", Detail: "scenario started"}); err != nil { + return err + } + var err error + switch scenario { + case "smoke": + err = h.runGoSuite(ctx, current, "e2e-client-agent", "real client/server WebSocket and control flow", "./test/integration/clientagent") + case "script": + err = h.runGoSuite(ctx, current, "e2e-script-transfer", "durable script transfer and replay", "./internal/client/agent") + case "recovery": + err = h.runStoreSuite(ctx, current) + default: + err = fmt.Errorf("unknown e2e scenario %q", scenario) + } + if err != nil { + return err + } + return h.appendJournal(current.RunID, journalEntry{At: h.now(), Step: "e2e-" + scenario, Status: "passed", Detail: "scenario passed"}) +} + func (h *harness) runStoreSuite(ctx context.Context, current *manifest) error { if err := h.appendJournal(current.RunID, journalEntry{At: h.now(), Step: "store-real-sqlite", Status: "running", Detail: "running real SQLite/WAL and filesystem cases"}); err != nil { return err diff --git a/test/harness/harness_test.go b/test/harness/harness_test.go index 00a728a..f27daa3 100644 --- a/test/harness/harness_test.go +++ b/test/harness/harness_test.go @@ -112,6 +112,17 @@ func TestIntegrationResumeValidation_HP_CFG_01(t *testing.T) { } } +func TestE2EScenarioValidationAndRunIdentity_HP_E2E_01(t *testing.T) { + t.Parallel() + h := &harness{root: t.TempDir(), now: time.Now, out: &bytes.Buffer{}} + if err := h.e2e(context.Background(), []string{"--scenario", "unknown"}); err == nil { + t.Fatal("unknown e2e scenario accepted") + } + if err := h.e2e(context.Background(), []string{"--resume", "--scenario", "smoke"}); err == nil { + t.Fatal("e2e resume without run ID accepted") + } +} + func TestBoundedSuiteLogAndFailedRunRecovery_BH_STORE_03(t *testing.T) { t.Parallel()