feat: execute durable client commands through supervisor
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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):
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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 {
|
||||
|
||||
@@ -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);
|
||||
`
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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[:]
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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) }
|
||||
@@ -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 }
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user