From ad564de23f2f102f118e24c9c854d9e94f67a1a0 Mon Sep 17 00:00:00 2001 From: cabbage Date: Sun, 6 Sep 2026 10:40:23 +0000 Subject: [PATCH] feat: persist and dispatch client command input --- internal/client/agent/dispatch.go | 7 +- internal/client/agent/handshake_test.go | 4 +- internal/client/spool/events.go | 104 +++++++++++++++- internal/client/spool/migrations.go | 21 +++- internal/client/spool/recovery.go | 35 +++++- internal/client/spool/spool.go | 34 ++++-- internal/client/spool/spool_test.go | 64 ++++++++++ internal/client/supervisor/supervisor.go | 101 ++++++++++++++++ internal/client/supervisor/supervisor_test.go | 32 +++++ .../client/supervisor/windows/environment.go | 114 ++++++++++++++++++ .../supervisor/windows/environment_test.go | 36 ++++++ internal/server/session/agent_server.go | 68 ++++++++++- internal/server/store/event.go | 13 ++ internal/server/store/stdin.go | 80 ++++++++++++ test/coverage.toml | 36 ++++++ .../clientagent_integration_test.go | 109 +++++++++++++++++ 16 files changed, 836 insertions(+), 22 deletions(-) create mode 100644 internal/client/supervisor/supervisor.go create mode 100644 internal/client/supervisor/supervisor_test.go create mode 100644 internal/client/supervisor/windows/environment.go create mode 100644 internal/client/supervisor/windows/environment_test.go diff --git a/internal/client/agent/dispatch.go b/internal/client/agent/dispatch.go index 44f7a03..dbfaa52 100644 --- a/internal/client/agent/dispatch.go +++ b/internal/client/agent/dispatch.go @@ -9,6 +9,7 @@ import ( "github.com/rvbox/rvbox/internal/agentproto" "github.com/rvbox/rvbox/internal/client/spool" "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/proto" ) // PersistDispatch admits a server dispatch to the durable spool before any @@ -30,5 +31,9 @@ func PersistDispatch(ctx context.Context, store *spool.Store, session Session, d if immutable == [32]byte{} { return spool.Acceptance{}, errors.New("empty dispatch immutable hash") } - return store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: immutable, Revision: dispatch.GetCommandRevision(), Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED)}, now) + spec, err := proto.MarshalOptions{Deterministic: true}.Marshal(dispatch.GetSpec()) + if err != nil { + return spool.Acceptance{}, err + } + return store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: immutable, Revision: dispatch.GetCommandRevision(), Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED), ExecutionSpec: spec}, now) } diff --git a/internal/client/agent/handshake_test.go b/internal/client/agent/handshake_test.go index 3fa0838..fbbbce2 100644 --- a/internal/client/agent/handshake_test.go +++ b/internal/client/agent/handshake_test.go @@ -1,6 +1,7 @@ package agent import ( + "bytes" "context" "errors" "path/filepath" @@ -92,7 +93,8 @@ func TestPersistDispatchUsesImmutableHash_HP_DISPATCH_04(t *testing.T) { defer store.Close() dispatch := &rvboxv1.CommandDispatch{IssueUuid: "019c46f1-1d02-7000-8000-000000000064", CommandRevision: 1, TargetSessionGeneration: 7, IssueTime: timestamppb.Now(), ImmutableRequestSha256: []byte("12345678901234567890123456789012"), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_POWERSHELL, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "Write-Output ok"}}} accepted, err := PersistDispatch(context.Background(), store, Session{Generation: 7}, dispatch, time.Now(), agentproto.DefaultLimits()) - if err != nil || accepted.Duplicate { + encodedSpec, marshalErr := proto.MarshalOptions{Deterministic: true}.Marshal(dispatch.GetSpec()) + if err != nil || marshalErr != nil || accepted.Duplicate || !bytes.Equal(accepted.Command.ExecutionSpec, encodedSpec) { t.Fatalf("PersistDispatch = %#v, %v", accepted, err) } duplicate, err := PersistDispatch(context.Background(), store, Session{Generation: 7}, dispatch, time.Now(), agentproto.DefaultLimits()) diff --git a/internal/client/spool/events.go b/internal/client/spool/events.go index aeb4c6a..7640b3a 100644 --- a/internal/client/spool/events.go +++ b/internal/client/spool/events.go @@ -1,12 +1,15 @@ package spool import ( + "bytes" "context" + "crypto/sha256" "database/sql" "errors" "fmt" "time" + "github.com/klauspost/compress/zstd" "github.com/rvbox/rvbox/internal/domain" ) @@ -16,6 +19,10 @@ type Command struct { Revision uint64 Phase uint32 Terminal bool + // ExecutionSpec is the deterministic protobuf payload received from the + // server. It is retained as raw protobuf bytes so the runtime can validate + // and execute exactly the admitted request after a restart. + ExecutionSpec []byte } type Acceptance struct { @@ -55,12 +62,28 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted if !validUUID(command.IssueUUID) || command.Revision == 0 || command.Phase == 0 || command.Phase > 11 || command.Terminal != isTerminalPhase(command.Phase) || acceptedAt.IsZero() { return Acceptance{}, errors.New("invalid command acceptance") } + var storedSpec []byte + var specCharge uint64 + if len(command.ExecutionSpec) > 0 { + if uint64(len(command.ExecutionSpec)) > store.maxExecutionSpecBytes { + return Acceptance{}, errors.New("execution specification exceeds client limit") + } + var err error + storedSpec, err = compressExecutionSpec(command.ExecutionSpec) + if err != nil { + return Acceptance{}, err + } + specCharge, err = EstimateCharge(ChargeInput{EncodedBytes: uint64(len(storedSpec)), SQLiteRows: 1, IndexEntries: 1}) + if err != nil { + return Acceptance{}, err + } + } tx, err := store.db.BeginTx(ctx, nil) if err != nil { return Acceptance{}, err } defer tx.Rollback() - existing, found, err := commandByUUID(ctx, tx, command.IssueUUID) + existing, found, err := commandByUUID(ctx, tx, command.IssueUUID, store.maxExecutionSpecBytes) if err != nil { return Acceptance{}, err } @@ -68,6 +91,9 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted if existing.ImmutableSHA256 != command.ImmutableSHA256 { return Acceptance{}, ErrCommandConflict } + if len(command.ExecutionSpec) > 0 && !bytes.Equal(existing.ExecutionSpec, command.ExecutionSpec) { + return Acceptance{}, ErrCommandConflict + } return Acceptance{Duplicate: true, Command: existing}, nil } var tombstoneHash []byte @@ -81,7 +107,7 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted if err != sql.ErrNoRows { return Acceptance{}, err } - charge, err := EstimateCharge(ChargeInput{SQLiteRows: 1, IndexEntries: 2}) + baseCharge, err := EstimateCharge(ChargeInput{SQLiteRows: 1, IndexEntries: 2}) if err != nil { return Acceptance{}, err } @@ -89,23 +115,51 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted if err != nil { return Acceptance{}, err } + charge, overflow := addChecked(baseCharge, specCharge) + if overflow { + return Acceptance{}, &CapacityError{Tier: CapacityTierHardMaximum, Requested: ^uint64(0), Available: store.quotaLimits.HardAllocationBytes} + } decision, err := CheckReservation(store.quotaLimits, ReservationState{ClientTotalCharged: clientTotal, CloseoutRemaining: store.quotaLimits.CloseoutReserveBytes}, ReservationRequest{ChargedBytes: charge}) if err != nil { return Acceptance{}, err } - _, err = tx.ExecContext(ctx, `INSERT INTO commands(issue_uuid, immutable_sha256, command_revision, phase, terminal, base_charged_bytes, total_charged_bytes, closeout_remaining_bytes, accepted_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, command.IssueUUID[:], command.ImmutableSHA256[:], command.Revision, command.Phase, boolInt(command.Terminal), charge, decision.CommandTotalCharged, decision.CloseoutRemaining, acceptedAt.UnixNano()) + _, err = tx.ExecContext(ctx, `INSERT INTO commands(issue_uuid, immutable_sha256, command_revision, phase, terminal, base_charged_bytes, total_charged_bytes, closeout_remaining_bytes, accepted_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, command.IssueUUID[:], command.ImmutableSHA256[:], command.Revision, command.Phase, boolInt(command.Terminal), baseCharge, decision.CommandTotalCharged, decision.CloseoutRemaining, acceptedAt.UnixNano()) if err != nil { return Acceptance{}, err } if err := updateClientTotalCharge(ctx, tx, decision.ClientTotalCharged); err != nil { return Acceptance{}, err } + if len(storedSpec) > 0 { + digest := immutableDigest(storedSpec) + if _, err := tx.ExecContext(ctx, `INSERT INTO command_specs(issue_uuid, raw_bytes, stored_bytes, compression, payload, payload_sha256, charged_bytes) VALUES (?, ?, ?, 2, ?, ?, ?)`, command.IssueUUID[:], len(command.ExecutionSpec), len(storedSpec), storedSpec, digest[:], specCharge); err != nil { + return Acceptance{}, err + } + } if err := tx.Commit(); err != nil { return Acceptance{}, err } return Acceptance{Command: command}, 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. +func (store *Store) GetCommand(ctx context.Context, issueUUID domain.UUID) (Command, error) { + if !validUUID(issueUUID) { + return Command{}, ErrUnknownCommand + } + command, found, err := commandByUUID(ctx, store.db, issueUUID, store.maxExecutionSpecBytes) + if err != nil { + return Command{}, err + } + if !found { + return Command{}, ErrUnknownCommand + } + command.ExecutionSpec = bytes.Clone(command.ExecutionSpec) + return command, nil +} + // AppendEvent assigns command-local order only. EventSeq remains zero until // AssignSendWindow makes the durable event eligible for a wire send. func (store *Store) AppendEvent(ctx context.Context, issueUUID domain.UUID, input EventInput) (Event, error) { @@ -448,11 +502,15 @@ func eventsBySequence(ctx context.Context, query queryer, issueUUID domain.UUID, func commandByUUID(ctx context.Context, query interface { QueryRowContext(context.Context, string, ...any) *sql.Row -}, issueUUID domain.UUID) (Command, bool, error) { +}, issueUUID domain.UUID, maxSpecBytes uint64) (Command, bool, error) { var command Command var hash []byte var terminal int - err := query.QueryRowContext(ctx, `SELECT immutable_sha256, command_revision, phase, terminal FROM commands WHERE issue_uuid = ?`, issueUUID[:]).Scan(&hash, &command.Revision, &command.Phase, &terminal) + var stored, digest []byte + var rawBytes, storedBytes, compression sql.NullInt64 + err := query.QueryRowContext(ctx, `SELECT c.immutable_sha256, c.command_revision, c.phase, c.terminal, + s.raw_bytes, s.stored_bytes, s.compression, s.payload, s.payload_sha256 + FROM commands c LEFT JOIN command_specs s ON s.issue_uuid = c.issue_uuid WHERE c.issue_uuid = ?`, issueUUID[:]).Scan(&hash, &command.Revision, &command.Phase, &terminal, &rawBytes, &storedBytes, &compression, &stored, &digest) if err == sql.ErrNoRows { return Command{}, false, nil } @@ -465,9 +523,45 @@ func commandByUUID(ctx context.Context, query interface { copy(command.ImmutableSHA256[:], hash) command.IssueUUID = issueUUID command.Terminal = terminal != 0 + if rawBytes.Valid { + computed := immutableDigest(stored) + if rawBytes.Int64 <= 0 || storedBytes.Int64 <= 0 || compression.Int64 != 2 || len(digest) != sha256.Size || uint64(rawBytes.Int64) > maxSpecBytes || uint64(storedBytes.Int64) != uint64(len(stored)) || !bytes.Equal(computed[:], digest) { + return Command{}, false, fmt.Errorf("invalid stored execution specification") + } + decoded, err := decompressExecutionSpec(stored, uint64(rawBytes.Int64), maxSpecBytes) + if err != nil { + return Command{}, false, err + } + command.ExecutionSpec = decoded + } return command, true, nil } +func compressExecutionSpec(raw []byte) ([]byte, error) { + encoder, err := zstd.NewWriter(nil, zstd.WithEncoderConcurrency(1)) + if err != nil { + return nil, err + } + defer encoder.Close() + return encoder.EncodeAll(raw, nil), nil +} + +func decompressExecutionSpec(stored []byte, rawBytes, maximum uint64) ([]byte, error) { + if rawBytes == 0 || rawBytes > maximum { + return nil, errors.New("invalid stored execution specification size") + } + decoder, err := zstd.NewReader(bytes.NewReader(stored), zstd.WithDecoderConcurrency(1), zstd.WithDecoderMaxMemory(maximum+1)) + if err != nil { + return nil, errors.New("invalid stored execution specification") + } + defer decoder.Close() + decoded, err := decoder.DecodeAll(stored, nil) + if err != nil || uint64(len(decoded)) != rawBytes || uint64(len(decoded)) > maximum { + return nil, errors.New("invalid stored execution specification") + } + return decoded, nil +} + func tombstonesToTrim(ctx context.Context, tx *sql.Tx, limit uint64) (uint64, error) { var count uint64 if err := tx.QueryRowContext(ctx, `SELECT count(*) FROM command_tombstones`).Scan(&count); err != nil { diff --git a/internal/client/spool/migrations.go b/internal/client/spool/migrations.go index 3900bee..8647322 100644 --- a/internal/client/spool/migrations.go +++ b/internal/client/spool/migrations.go @@ -13,7 +13,10 @@ type migration struct { sql string } -var migrations = []migration{{version: 1, sql: schemaV1}} +var migrations = []migration{ + {version: 1, sql: schemaV1}, + {version: 2, sql: schemaV2}, +} func applyMigrations(ctx context.Context, db *sql.DB) error { if _, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations ( @@ -113,3 +116,19 @@ CREATE TABLE command_tombstones ( ) STRICT; CREATE INDEX tombstones_fifo ON command_tombstones(acknowledged_at, issue_uuid); ` + +// schemaV2 adds the immutable execution specification retained with every +// accepted dispatch. Keeping it in its own table lets old commands created by +// early development builds (which have no specification) remain readable and +// makes the payload's quota charge independently auditable. +const schemaV2 = ` +CREATE TABLE command_specs ( + issue_uuid BLOB PRIMARY KEY REFERENCES commands(issue_uuid) ON DELETE CASCADE CHECK(length(issue_uuid) = 16), + raw_bytes INTEGER NOT NULL CHECK(raw_bytes > 0), + stored_bytes INTEGER NOT NULL CHECK(stored_bytes > 0), + compression INTEGER NOT NULL CHECK(compression = 2), + payload BLOB NOT NULL, + payload_sha256 BLOB NOT NULL CHECK(length(payload_sha256) = 32), + charged_bytes INTEGER NOT NULL CHECK(charged_bytes > 0) +) STRICT, WITHOUT ROWID; +` diff --git a/internal/client/spool/recovery.go b/internal/client/spool/recovery.go index 22970a4..b7d450f 100644 --- a/internal/client/spool/recovery.go +++ b/internal/client/spool/recovery.go @@ -19,6 +19,7 @@ var ( type RecoveryReport struct { EventsChecked uint64 ScriptsChecked uint64 + SpecsChecked uint64 } // Check verifies state that SQLite's own structural check cannot express. The @@ -34,7 +35,7 @@ func (store *Store) Check(ctx context.Context) (RecoveryReport, error) { if err := checkQuotaCounters(ctx, store.db); err != nil { return RecoveryReport{}, err } - return checkPayloads(ctx, store.db, store.maxScriptBytes) + return checkPayloads(ctx, store.db, store.maxScriptBytes, store.maxExecutionSpecBytes) } func quickCheck(ctx context.Context, database *sql.DB) error { @@ -65,6 +66,7 @@ func checkQuotaCounters(ctx context.Context, database *sql.DB) error { err := database.QueryRowContext(ctx, `SELECT count(*) FROM commands AS c WHERE c.total_charged_bytes != c.base_charged_bytes + COALESCE((SELECT sum(charged_bytes) FROM events WHERE issue_uuid = c.issue_uuid), 0) + + COALESCE((SELECT charged_bytes FROM command_specs WHERE issue_uuid = c.issue_uuid), 0) + COALESCE((SELECT charged_bytes FROM scripts WHERE issue_uuid = c.issue_uuid), 0) OR c.output_charged_bytes != COALESCE((SELECT sum(charged_bytes) FROM events WHERE issue_uuid = c.issue_uuid AND output = 1), 0)`).Scan(&invalid) @@ -81,7 +83,7 @@ OR c.output_charged_bytes != return nil } -func checkPayloads(ctx context.Context, database *sql.DB, maxScriptBytes uint64) (RecoveryReport, error) { +func checkPayloads(ctx context.Context, database *sql.DB, maxScriptBytes, maxExecutionSpecBytes uint64) (RecoveryReport, error) { var report RecoveryReport rows, err := database.QueryContext(ctx, `SELECT issue_uuid, event_kind, raw_bytes, payload, payload_sha256 FROM events ORDER BY issue_uuid, local_ordinal`) if err != nil { @@ -119,6 +121,35 @@ func checkPayloads(ctx context.Context, database *sql.DB, maxScriptBytes uint64) if err := rows.Close(); err != nil { return report, err } + specs, err := database.QueryContext(ctx, `SELECT issue_uuid, raw_bytes, stored_bytes, compression, payload, payload_sha256 FROM command_specs ORDER BY issue_uuid`) + if err != nil { + return report, err + } + for specs.Next() { + var owner, payload, digest []byte + var rawBytes, storedBytes uint64 + var compression uint32 + if err := specs.Scan(&owner, &rawBytes, &storedBytes, &compression, &payload, &digest); err != nil { + _ = specs.Close() + return report, err + } + computed := immutableDigest(payload) + if _, valid := copiedUUID(owner); !valid || compression != 2 || rawBytes == 0 || rawBytes > maxExecutionSpecBytes || storedBytes != uint64(len(payload)) || len(digest) != 32 || !bytes.Equal(computed[:], digest) { + _ = specs.Close() + return report, ErrStoredPayloadChecksum + } + if _, err := decompressExecutionSpec(payload, rawBytes, maxExecutionSpecBytes); err != nil { + _ = specs.Close() + return report, ErrStoredPayloadChecksum + } + report.SpecsChecked++ + } + if err := specs.Close(); err != nil { + return report, err + } + if err := specs.Err(); err != nil { + return report, err + } scripts, err := database.QueryContext(ctx, `SELECT issue_uuid, received_raw_bytes, stored_bytes, stored_data FROM scripts ORDER BY issue_uuid`) if err != nil { return report, err diff --git a/internal/client/spool/spool.go b/internal/client/spool/spool.go index 33b21b1..31ee03a 100644 --- a/internal/client/spool/spool.go +++ b/internal/client/spool/spool.go @@ -37,18 +37,23 @@ type Options struct { BusyTimeout time.Duration TombstoneLimit uint64 QuotaLimits QuotaLimits - MaxScriptBytes uint64 + // MaxExecutionSpecBytes bounds the decompressed protobuf retained for an + // accepted command. It is separate from MaxScriptBytes because a script + // body is uploaded in a different table and lifecycle. + MaxExecutionSpecBytes uint64 + MaxScriptBytes uint64 } type Store struct { - db *sql.DB - unlock func() error - dataDir string - identity domain.UUID - tombstoneLimit uint64 - quotaLimits QuotaLimits - maxScriptBytes uint64 - mu sync.Mutex + db *sql.DB + unlock func() error + dataDir string + identity domain.UUID + tombstoneLimit uint64 + quotaLimits QuotaLimits + maxScriptBytes uint64 + maxExecutionSpecBytes uint64 + mu sync.Mutex } // Open recovers an existing spool or creates an empty one. The caller must @@ -76,6 +81,15 @@ func Open(ctx context.Context, options Options) (*Store, error) { options.MaxScriptBytes = options.QuotaLimits.CommandTotalBytes } } + if options.MaxExecutionSpecBytes == 0 { + options.MaxExecutionSpecBytes = 768 << 10 + if options.MaxExecutionSpecBytes > options.QuotaLimits.CommandTotalBytes { + options.MaxExecutionSpecBytes = options.QuotaLimits.CommandTotalBytes + } + } + if options.MaxExecutionSpecBytes > options.QuotaLimits.CommandTotalBytes { + return nil, errors.New("execution specification maximum exceeds client command quota") + } if options.MaxScriptBytes > options.QuotaLimits.CommandTotalBytes { return nil, errors.New("script maximum exceeds client command quota") } @@ -110,7 +124,7 @@ func Open(ctx context.Context, options Options) (*Store, error) { } db.SetMaxOpenConns(1) db.SetMaxIdleConns(1) - store := &Store{db: db, unlock: unlock, dataDir: options.DataDir, identity: identity, tombstoneLimit: options.TombstoneLimit, quotaLimits: options.QuotaLimits, maxScriptBytes: options.MaxScriptBytes} + store := &Store{db: db, unlock: unlock, dataDir: options.DataDir, identity: identity, tombstoneLimit: options.TombstoneLimit, quotaLimits: options.QuotaLimits, maxScriptBytes: options.MaxScriptBytes, maxExecutionSpecBytes: options.MaxExecutionSpecBytes} if err := db.PingContext(ctx); err != nil { _ = store.Close() return nil, fmt.Errorf("open client spool SQLite: %w", err) diff --git a/internal/client/spool/spool_test.go b/internal/client/spool/spool_test.go index 2de09f4..c8e176f 100644 --- a/internal/client/spool/spool_test.go +++ b/internal/client/spool/spool_test.go @@ -1,6 +1,7 @@ package spool import ( + "bytes" "context" "crypto/sha256" "errors" @@ -12,6 +13,69 @@ import ( "github.com/rvbox/rvbox/internal/domain" ) +func TestExecutionSpecIsDurableAndQuotaCounted_HP_DISPATCH_07(t *testing.T) { + t.Parallel() + ctx := context.Background() + directory := filepath.Join(t.TempDir(), "spool") + store := openTestStore(t, ctx, directory, DefaultTombstoneLimit) + issue := testUUID(t, "019c46f1-1d02-7000-8000-000000000002") + spec := []byte("deterministic execution specification") + command := testCommand(issue, []byte("spec-hash")) + command.ExecutionSpec = spec + acceptedAt := time.Date(2026, time.September, 6, 12, 0, 0, 0, time.UTC) + accepted, err := store.AcceptCommand(ctx, command, acceptedAt) + if err != nil || accepted.Duplicate { + t.Fatalf("spec acceptance = %#v, %v", accepted, err) + } + if !bytes.Equal(accepted.Command.ExecutionSpec, spec) { + t.Fatalf("accepted spec = %q, want %q", accepted.Command.ExecutionSpec, spec) + } + var base, total, specCharge uint64 + if err := store.db.QueryRowContext(ctx, `SELECT base_charged_bytes, total_charged_bytes FROM commands WHERE issue_uuid = ?`, issue[:]).Scan(&base, &total); err != nil { + t.Fatal(err) + } + if err := store.db.QueryRowContext(ctx, `SELECT charged_bytes FROM command_specs WHERE issue_uuid = ?`, issue[:]).Scan(&specCharge); err != nil { + t.Fatal(err) + } + if specCharge == 0 || total != base+specCharge { + t.Fatalf("spec charge accounting = base=%d total=%d spec=%d", base, total, specCharge) + } + duplicate, err := store.AcceptCommand(ctx, command, acceptedAt.Add(time.Second)) + if err != nil || !duplicate.Duplicate || !bytes.Equal(duplicate.Command.ExecutionSpec, spec) { + t.Fatalf("spec replay = %#v, %v", duplicate, err) + } + if err := store.Close(); err != nil { + t.Fatal(err) + } + store = openTestStore(t, ctx, directory, DefaultTombstoneLimit) + reopened, err := store.AcceptCommand(ctx, command, acceptedAt.Add(2*time.Second)) + if err != nil || !reopened.Duplicate || !bytes.Equal(reopened.Command.ExecutionSpec, spec) { + t.Fatalf("spec replay after reopen = %#v, %v", reopened, err) + } + report, err := store.Check(ctx) + if err != nil || report.SpecsChecked != 1 { + t.Fatalf("spec recovery report = %#v, %v", report, err) + } +} + +func TestExecutionSpecCorruptionMarksSpoolDirty_BH_DISPATCH_07(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-000000000003") + command := testCommand(issue, []byte("spec-corrupt")) + command.ExecutionSpec = []byte("payload") + if _, err := store.AcceptCommand(ctx, command, time.Now().UTC()); err != nil { + t.Fatal(err) + } + if _, err := store.db.ExecContext(ctx, `UPDATE command_specs SET payload = ? WHERE issue_uuid = ?`, []byte("tampered"), issue[:]); err != nil { + t.Fatal(err) + } + if _, err := store.Check(ctx); !errors.Is(err, ErrStoredPayloadChecksum) { + t.Fatalf("corrupt execution spec Check error = %v, want ErrStoredPayloadChecksum", err) + } +} + func TestStoreDurableIdentityAndSingleInstance_HP_CLIENT_01(t *testing.T) { t.Parallel() ctx := context.Background() diff --git a/internal/client/supervisor/supervisor.go b/internal/client/supervisor/supervisor.go new file mode 100644 index 0000000..46d7811 --- /dev/null +++ b/internal/client/supervisor/supervisor.go @@ -0,0 +1,101 @@ +// Package supervisor defines the narrow process-control seam shared by the +// client runtime and platform implementations. It intentionally contains no +// operating-system handles or process APIs. +package supervisor + +import ( + "context" + "errors" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/domain" +) + +var ( + ErrInvalidStartSpec = errors.New("invalid supervisor start specification") + ErrUnsupported = errors.New("supervisor operation is unsupported") +) + +type SignalKind uint8 + +const ( + SignalTerm SignalKind = iota + 1 + SignalKill +) + +// StartSpec is already validated at the protocol boundary. The supervisor +// 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 + WorkingDirectory string + Environment map[string]string + ExecutionProfiles []string +} + +func (spec StartSpec) Validate() error { + if _, err := domain.ParseUUIDv7(spec.IssueUUID.String()); err != nil || spec.CommandRevision == 0 || spec.Execution == nil || spec.Execution.Source == nil { + return ErrInvalidStartSpec + } + if spec.WorkingDirectory == "" { + return ErrInvalidStartSpec + } + return nil +} + +// EffectiveIdentity is immutable process evidence captured before launch. +// Empty user/session fields mean a Session 0 service context. +type EffectiveIdentity struct { + Context string + SessionID uint32 + UserSID string + LogonSID string + Elevated bool + Integrity string +} + +type Process interface { + IssueUUID() domain.UUID + Identity() EffectiveIdentity + Wait(context.Context) (ExitStatus, error) + WriteStdin(context.Context, []byte, bool) error + CloseStdin(context.Context) error +} + +type ExitStatus struct { + Code int32 + Signaled bool + Signal SignalKind + StartedAt time.Time + FinishedAt time.Time + OutputDrained bool + Output bool +} + +type SignalOutcome struct { + Delivered bool + Escalated bool + Detail string + ObservedAt time.Time +} + +type ResourceSnapshot struct { + CPUTime time.Duration + ResidentBytes uint64 + IOReadBytes uint64 + IOWriteBytes uint64 + ProcessCount uint64 + ObservedAt time.Time + Complete bool + Detail string +} + +type Supervisor interface { + Start(context.Context, StartSpec) (Process, error) + Signal(context.Context, Process, SignalKind) (SignalOutcome, error) + Snapshot(context.Context, Process) (ResourceSnapshot, error) + StopAll(context.Context) error +} diff --git a/internal/client/supervisor/supervisor_test.go b/internal/client/supervisor/supervisor_test.go new file mode 100644 index 0000000..4cd97f1 --- /dev/null +++ b/internal/client/supervisor/supervisor_test.go @@ -0,0 +1,32 @@ +package supervisor + +import ( + "errors" + "testing" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/domain" +) + +func TestStartSpecValidation_BH_SUPERVISOR_01(t *testing.T) { + t.Parallel() + issue, err := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000091") + if err != nil { + t.Fatal(err) + } + cases := []StartSpec{ + {}, + {IssueUUID: issue, CommandRevision: 1, WorkingDirectory: `C:\work`}, + {IssueUUID: issue, CommandRevision: 1, Execution: &rvboxv1.ExecutionSpec{}, WorkingDirectory: `C:\work`}, + {IssueUUID: issue, CommandRevision: 1, Execution: &rvboxv1.ExecutionSpec{Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo ok"}}}, + } + for index, spec := range cases { + if !errors.Is(spec.Validate(), ErrInvalidStartSpec) { + t.Fatalf("case %d validation = %v, want ErrInvalidStartSpec", index, spec.Validate()) + } + } + valid := StartSpec{IssueUUID: issue, CommandRevision: 1, Execution: &rvboxv1.ExecutionSpec{Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo ok"}}, WorkingDirectory: `C:\work`} + if err := valid.Validate(); err != nil { + t.Fatal(err) + } +} diff --git a/internal/client/supervisor/windows/environment.go b/internal/client/supervisor/windows/environment.go new file mode 100644 index 0000000..cef870d --- /dev/null +++ b/internal/client/supervisor/windows/environment.go @@ -0,0 +1,114 @@ +package windows + +import ( + "errors" + "sort" + "strings" + "unicode/utf16" + "unicode/utf8" +) + +var ( + ErrInvalidEnvironmentKey = errors.New("invalid Windows environment key") + ErrInvalidEnvironmentValue = errors.New("invalid Windows environment value") + ErrDuplicateEnvironmentKey = errors.New("duplicate Windows environment key") + ErrEnvironmentTooLarge = errors.New("Windows environment block is too large") +) + +// EnvironmentEntry is an ordered, case-preserving environment assignment. +// Windows compares names case-insensitively, so Folded is never emitted and is +// used only to make duplicate handling deterministic. +type EnvironmentEntry struct { + Key string + Value string + Folded string +} + +// MergeEnvironment overlays request values onto the token-derived base +// environment. Shell resolution happens before this function is called; an +// override therefore cannot redirect the configured executable. +func MergeEnvironment(base, overrides map[string]string) ([]EnvironmentEntry, error) { + entries := make(map[string]EnvironmentEntry, len(base)+len(overrides)) + add := func(key, value string, override bool) error { + if !validEnvironmentKey(key, override) { + return ErrInvalidEnvironmentKey + } + if strings.IndexByte(value, 0) >= 0 || !utf8.ValidString(value) { + return ErrInvalidEnvironmentValue + } + folded := strings.ToUpper(key) + if _, exists := entries[folded]; exists { + if !override { + return ErrDuplicateEnvironmentKey + } + delete(entries, folded) + } + entries[folded] = EnvironmentEntry{Key: key, Value: value, Folded: folded} + return nil + } + baseKeys := make([]string, 0, len(base)) + for key := range base { + baseKeys = append(baseKeys, key) + } + sort.Strings(baseKeys) + for _, key := range baseKeys { + if err := add(key, base[key], false); err != nil { + return nil, err + } + } + overrideKeys := make([]string, 0, len(overrides)) + for key := range overrides { + overrideKeys = append(overrideKeys, key) + } + sort.Strings(overrideKeys) + for _, key := range overrideKeys { + if err := add(key, overrides[key], true); err != nil { + return nil, err + } + } + result := make([]EnvironmentEntry, 0, len(entries)) + for _, entry := range entries { + result = append(result, entry) + } + sort.Slice(result, func(left, right int) bool { + if result[left].Folded == result[right].Folded { + return result[left].Key < result[right].Key + } + return result[left].Folded < result[right].Folded + }) + return result, nil +} + +// BuildEnvironmentBlock converts the merged entries to the UTF-16 block +// expected by CreateProcessAsUser. The final two NUL code units are included. +func BuildEnvironmentBlock(base, overrides map[string]string) ([]uint16, error) { + entries, err := MergeEnvironment(base, overrides) + if err != nil { + return nil, err + } + var units []uint16 + for _, entry := range entries { + units = append(units, utf16.Encode([]rune(entry.Key+"="+entry.Value))...) + units = append(units, 0) + if len(units) > 32767 { + return nil, ErrEnvironmentTooLarge + } + } + units = append(units, 0) + return units, nil +} + +func validEnvironmentKey(key string, override bool) bool { + if key == "" || strings.IndexByte(key, 0) >= 0 || !utf8.ValidString(key) { + return false + } + if override && (strings.ContainsRune(key, '=') || strings.HasPrefix(key, "=")) { + return false + } + if strings.HasPrefix(key, "=") { + // CreateEnvironmentBlock may return drive-current-directory pseudo + // variables such as =C:=C:\\work. They are accepted only as base data. + return !override && strings.Count(key, "=") == 1 + } + return !strings.ContainsRune(key, '=') +} diff --git a/internal/client/supervisor/windows/environment_test.go b/internal/client/supervisor/windows/environment_test.go new file mode 100644 index 0000000..654389c --- /dev/null +++ b/internal/client/supervisor/windows/environment_test.go @@ -0,0 +1,36 @@ +package windows + +import ( + "errors" + "testing" +) + +func TestEnvironmentOverlayIsCaseInsensitiveAndDeterministic_HP_WINENV_01(t *testing.T) { + t.Parallel() + base := map[string]string{"Path": `C:\Windows`, "TEMP": `C:\Temp`, "=C:": `C:\Windows`} + overrides := map[string]string{"PATH": `C:\Pinned`, "Name": "rvbox"} + entries, err := MergeEnvironment(base, overrides) + if err != nil { + t.Fatal(err) + } + if len(entries) != 4 || entries[0].Folded != "=C:" || entries[1].Folded != "NAME" || entries[2].Folded != "PATH" || entries[2].Value != `C:\Pinned` { + t.Fatalf("merged entries = %#v", entries) + } + block, err := BuildEnvironmentBlock(base, overrides) + if err != nil || len(block) == 0 || block[len(block)-1] != 0 || block[len(block)-2] != 0 { + t.Fatalf("environment block = len=%d err=%v", len(block), err) + } +} + +func TestEnvironmentOverlayRejectsMalformedAndDuplicateBase_BH_WINENV_01(t *testing.T) { + t.Parallel() + if _, err := MergeEnvironment(map[string]string{"Path": "one", "PATH": "two"}, nil); !errors.Is(err, ErrDuplicateEnvironmentKey) { + t.Fatalf("duplicate base error = %v", err) + } + if _, err := MergeEnvironment(nil, map[string]string{"BAD=KEY": "value"}); !errors.Is(err, ErrInvalidEnvironmentKey) { + t.Fatalf("invalid override key error = %v", err) + } + if _, err := MergeEnvironment(nil, map[string]string{"NUL": "bad\x00value"}); !errors.Is(err, ErrInvalidEnvironmentValue) { + t.Fatalf("invalid override value error = %v", err) + } +} diff --git a/internal/server/session/agent_server.go b/internal/server/session/agent_server.go index 6620fdb..0311dd4 100644 --- a/internal/server/session/agent_server.go +++ b/internal/server/session/agent_server.go @@ -141,13 +141,14 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w capacity := &CapacityShadow{} var capacityMu sync.Mutex reservations := make(map[string]DispatchLane) + stdinSent := make(map[string]struct{}) reconciled := make(chan struct{}) var reconcileOnce sync.Once if !capacity.UpdateAdvertised(0, 0, hello.GetMaxRunningCommands(), hello.GetMaxQueuedCommands()) { server.close(connection, websocket.StatusPolicyViolation, "invalid initial client capacity") return } - go server.dispatchLoop(sessionContext, queue, handle.DispatchWake(), reconciled, &capacityMu, capacity, reservations, hello.GetClientId(), hello.GetPlatform(), encodeSessionID(sessionID), registration.Generation, cancel) + go server.dispatchLoop(sessionContext, queue, handle.DispatchWake(), reconciled, &capacityMu, capacity, reservations, stdinSent, hello.GetClientId(), hello.GetPlatform(), encodeSessionID(sessionID), registration.Generation, cancel) encodedSessionID := encodeSessionID(sessionID) welcome, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: encodedSessionID, SessionGeneration: registration.Generation, @@ -265,6 +266,13 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w server.close(connection, websocket.StatusInternalError, "could not acknowledge command event") return } + if stdinAck := event.GetStdinAck(); stdinAck != nil { + issue, _ := domain.ParseUUIDv7(event.GetIssueUuid()) + capacityMu.Lock() + delete(stdinSent, stdinIntentKey(issue, stdinAck.GetWriteSeq())) + capacityMu.Unlock() + handle.SignalDispatch() + } } } } @@ -366,10 +374,58 @@ func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *Write return candidate.IssueUUID, true, nil } +// enqueueNextStdin exposes one durable input intent on the essential control +// lane. Session-local sent tracking suppresses duplicate frames while a live +// connection remains usable; reconnecting naturally replays unacknowledged +// writes from storage. +func (server *AgentServer) enqueueNextStdin(ctx context.Context, queue *WriterQueue, clientID, sessionID string, generation uint64, sent map[string]struct{}, sentMu *sync.Mutex) (bool, error) { + intents, err := server.Store.PendingStdin(ctx, clientID) + if err != nil { + return false, err + } + for _, intent := range intents { + key := stdinIntentKey(intent.IssueUUID, intent.WriteSeq) + sentMu.Lock() + _, alreadySent := sent[key] + if !alreadySent { + sent[key] = struct{}{} + } + sentMu.Unlock() + if alreadySent { + continue + } + envelope := &rvboxv1.AgentEnvelope{SessionId: sessionID, SessionGeneration: generation} + if intent.Close { + envelope.Payload = &rvboxv1.AgentEnvelope_CloseStdin{CloseStdin: &rvboxv1.CloseStdin{IssueUuid: intent.IssueUUID.String(), WriteSeq: intent.WriteSeq}} + } else { + envelope.Payload = &rvboxv1.AgentEnvelope_StdinWrite{StdinWrite: &rvboxv1.StdinWrite{IssueUuid: intent.IssueUUID.String(), WriteSeq: intent.WriteSeq, Data: append([]byte(nil), intent.Data...), AppendNewline: intent.AppendNewline}} + } + encoded, marshalErr := proto.Marshal(envelope) + if marshalErr != nil { + sentMu.Lock() + delete(sent, key) + sentMu.Unlock() + return false, marshalErr + } + if enqueueErr := queue.EnqueueControl(Frame{Kind: FrameControl, Payload: encoded}); enqueueErr != nil { + sentMu.Lock() + delete(sent, key) + sentMu.Unlock() + return false, enqueueErr + } + return true, nil + } + return false, nil +} + +func stdinIntentKey(issue domain.UUID, writeSeq uint64) string { + return issue.String() + ":" + fmt.Sprint(writeSeq) +} + // dispatchLoop is the per-session serialized dispatcher. It waits for a // complete reconciliation result before consuming queued work, then coalesces // wakeups from local control RPCs, capacity advertisements, and acceptances. -func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue, wake <-chan struct{}, reconciled <-chan struct{}, capacityMu *sync.Mutex, capacity *CapacityShadow, reservations map[string]DispatchLane, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, cancel context.CancelFunc) { +func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue, wake <-chan struct{}, reconciled <-chan struct{}, capacityMu *sync.Mutex, capacity *CapacityShadow, reservations map[string]DispatchLane, stdinSent map[string]struct{}, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, cancel context.CancelFunc) { select { case <-reconciled: case <-ctx.Done(): @@ -377,6 +433,14 @@ func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue, } for { for { + stdinQueued, stdinErr := server.enqueueNextStdin(ctx, queue, clientID, sessionID, generation, stdinSent, capacityMu) + if stdinErr != nil { + server.closeForDispatchFailure(cancel) + return + } + if stdinQueued { + continue + } capacityMu.Lock() lane := capacity.Reserve() capacityMu.Unlock() diff --git a/internal/server/store/event.go b/internal/server/store/event.go index a3473f6..994d0bb 100644 --- a/internal/server/store/event.go +++ b/internal/server/store/event.go @@ -12,6 +12,7 @@ import ( rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/domain" + "google.golang.org/protobuf/proto" ) const ( @@ -235,6 +236,18 @@ retained_compressed_bytes = retained_compressed_bytes + ?, output_charged_bytes } } } + if err == nil && event.EventType == 7 { + // Stdin acknowledgements are command events, but their durable delivery + // cursor lives in stdin_writes. Mark the cumulative prefix in the same + // transaction as the event so a crash cannot make an acknowledged write + // replay forever or release it before its event is durable. + var wire rvboxv1.CommandEvent + if unmarshalErr := proto.Unmarshal(event.Payload, &wire); unmarshalErr != nil || wire.GetStdinAck() == nil || wire.GetStdinAck().GetWriteSeq() == 0 { + err = ErrInvalidSegmentRecord + } else { + _, err = tx.ExecContext(ctx, `UPDATE stdin_writes SET acknowledged = 1 WHERE issue_uuid = ? AND write_seq <= ?`, event.IssueUUID[:], wire.GetStdinAck().GetWriteSeq()) + } + } if err == nil { var update sql.Result update, err = tx.ExecContext(ctx, `UPDATE clients SET charged_bytes = ? WHERE client_id = ? AND charged_bytes = ?`, reservation.ClientTotalCharged, clientID, clientCharged) diff --git a/internal/server/store/stdin.go b/internal/server/store/stdin.go index a773f9e..3d24d30 100644 --- a/internal/server/store/stdin.go +++ b/internal/server/store/stdin.go @@ -1,6 +1,7 @@ package store import ( + "bytes" "context" "crypto/sha256" "database/sql" @@ -9,6 +10,7 @@ import ( "math" "time" + "github.com/klauspost/compress/zstd" "github.com/rvbox/rvbox/internal/domain" ) @@ -33,6 +35,84 @@ type StdinWriteResult struct { Duplicate bool } +// StdinIntent is a server-owned input record ready for delivery to a live +// client session. It is returned in command/write order and contains a copy of +// the decompressed payload; callers must not mutate storage-owned buffers. +type StdinIntent struct { + IssueUUID domain.UUID + WriteSeq uint64 + Data []byte + AppendNewline bool + Close bool +} + +// PendingStdin returns unacknowledged input intents for one client. Delivery is +// intentionally tracked by the session dispatcher, not by this durable query: +// a lost connection simply causes the next session to replay the same write +// sequence, which the client acknowledges idempotently. +func (store *Store) PendingStdin(ctx context.Context, clientID string) ([]StdinIntent, error) { + if clientID == "" { + return nil, errors.New("client ID is required") + } + database, err := store.openDatabase() + if err != nil { + return nil, err + } + rows, err := database.QueryContext(ctx, `SELECT w.issue_uuid, w.write_seq, w.payload, w.raw_bytes, w.stored_bytes, w.compression, w.sha256, w.append_newline, w.close_intent +FROM stdin_writes w JOIN commands c ON c.issue_uuid = w.issue_uuid +WHERE c.client_id = ? AND c.lifecycle BETWEEN 1 AND 4 AND w.acknowledged = 0 +ORDER BY c.issue_time, c.issue_uuid, w.write_seq`, clientID) + if err != nil { + return nil, err + } + defer rows.Close() + var result []StdinIntent + for rows.Next() { + var encodedIssue, stored, digest []byte + var writeSeq, rawBytes, storedBytes uint64 + var compression uint32 + var appendNewline, closeIntent bool + if err := rows.Scan(&encodedIssue, &writeSeq, &stored, &rawBytes, &storedBytes, &compression, &digest, &appendNewline, &closeIntent); err != nil { + return nil, err + } + if len(encodedIssue) != len(domain.UUID{}) || writeSeq == 0 || compression != 2 || storedBytes != uint64(len(stored)) || len(digest) != sha256.Size || rawBytes > maxStoredExecutionSpecBytes { + return nil, ErrInvalidSegmentRecord + } + var data []byte + data, err = decompressStdinPayload(stored, rawBytes) + if err != nil || (!closeIntent && rawBytes == 0) || (closeIntent && rawBytes != 0) { + return nil, ErrInvalidSegmentRecord + } + computed := sha256.Sum256(data) + if string(computed[:]) != string(digest) { + return nil, ErrInvalidSegmentRecord + } + var issue domain.UUID + copy(issue[:], encodedIssue) + result = append(result, StdinIntent{IssueUUID: issue, WriteSeq: writeSeq, Data: data, AppendNewline: appendNewline, Close: closeIntent}) + } + if err := rows.Err(); err != nil { + return nil, err + } + return result, nil +} + +func decompressStdinPayload(stored []byte, expected uint64) ([]byte, error) { + if expected > maxStoredExecutionSpecBytes { + return nil, ErrInvalidSegmentRecord + } + decoder, err := zstd.NewReader(bytes.NewReader(stored), zstd.WithDecoderConcurrency(1), zstd.WithDecoderMaxMemory(maxStoredExecutionSpecBytes+1)) + if err != nil { + return nil, ErrInvalidSegmentRecord + } + defer decoder.Close() + data, err := decoder.DecodeAll(stored, nil) + if err != nil || uint64(len(data)) != expected { + return nil, ErrInvalidSegmentRecord + } + return data, nil +} + // AppendStdin durably records one ordered stdin intent. It does not claim // delivery to a process; an agent acknowledgement is a later protocol event. // Reusing RequestUUID with the same immutable hash returns the original write diff --git a/test/coverage.toml b/test/coverage.toml index 25106ac..2478e3c 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -227,6 +227,24 @@ layer = "unit" status = "implemented" tests = ["internal/client/supervisor/windows/selection_test.go:TestSelectExecutionContextBoundaries_BH_WINCTX_02"] +[[requirements]] +id = "HP-WINENV-01" +layer = "unit" +status = "implemented" +tests = ["internal/client/supervisor/windows/environment_test.go:TestEnvironmentOverlayIsCaseInsensitiveAndDeterministic_HP_WINENV_01"] + +[[requirements]] +id = "BH-WINENV-01" +layer = "unit" +status = "implemented" +tests = ["internal/client/supervisor/windows/environment_test.go:TestEnvironmentOverlayRejectsMalformedAndDuplicateBase_BH_WINENV_01"] + +[[requirements]] +id = "BH-SUPERVISOR-01" +layer = "unit" +status = "implemented" +tests = ["internal/client/supervisor/supervisor_test.go:TestStartSpecValidation_BH_SUPERVISOR_01"] + [[requirements]] id = "BH-SES-01" layer = "unit" @@ -302,6 +320,24 @@ layer = "integration" status = "implemented" tests = ["test/integration/clientagent/clientagent_integration_test.go:TestControlQueueWakesReconciledSession_HP_DISPATCH_06"] +[[requirements]] +id = "HP-DISPATCH-07" +layer = "unit" +status = "implemented" +tests = ["internal/client/spool/spool_test.go:TestExecutionSpecIsDurableAndQuotaCounted_HP_DISPATCH_07"] + +[[requirements]] +id = "BH-DISPATCH-07" +layer = "unit" +status = "implemented" +tests = ["internal/client/spool/spool_test.go:TestExecutionSpecCorruptionMarksSpoolDirty_BH_DISPATCH_07"] + +[[requirements]] +id = "HP-DISPATCH-08" +layer = "integration" +status = "implemented" +tests = ["test/integration/clientagent/clientagent_integration_test.go:TestControlStdinIntentReplaysAndAcknowledges_HP_DISPATCH_08"] + [[requirements]] id = "HP-EVENT-01" layer = "unit" diff --git a/test/integration/clientagent/clientagent_integration_test.go b/test/integration/clientagent/clientagent_integration_test.go index b066580..9b33143 100644 --- a/test/integration/clientagent/clientagent_integration_test.go +++ b/test/integration/clientagent/clientagent_integration_test.go @@ -193,6 +193,115 @@ func TestControlQueueWakesReconciledSession_HP_DISPATCH_06(t *testing.T) { } } +func TestControlStdinIntentReplaysAndAcknowledges_HP_DISPATCH_08(t *testing.T) { + ctx := context.Background() + persistence, err := store.Open(ctx, store.Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer persistence.Close() + registry := session.NewRegistry() + server := httptest.NewServer(&session.AgentServer{Store: persistence, Registry: registry, Path: "/v1/agent", Limits: agentproto.DefaultLimits()}) + defer server.Close() + controlService, err := control.NewService(control.Options{Store: persistence, WakeClient: registry.Wake, CursorKey: []byte("0123456789abcdef0123456789abcdef")}) + if err != nil { + t.Fatal(err) + } + address := "ws" + strings.TrimPrefix(server.URL, "http") + "/v1/agent" + transport, err := agent.DialWebSocket(ctx, address, nil) + if err != nil { + t.Fatal(err) + } + defer transport.Close() + hello := &rvboxv1.ClientHello{ClientId: "win-stdin-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: "019c46f1-1d02-7000-8000-000000000079", MaxRunningCommands: 1, MaxQueuedCommands: 1, SentAt: timestamppb.Now()} + accepted, err := agent.Handshake(ctx, transport, hello, agentproto.DefaultLimits()) + if err != nil { + t.Fatal(err) + } + issue := fixedIssueID(0x7a) + spec := &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_POWERSHELL, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "Read-Host"}} + serialized, err := proto.Marshal(spec) + if err != nil { + t.Fatal(err) + } + hash := sha256.Sum256([]byte("stdin request")) + parsedIssue, _ := domain.ParseUUIDv7(issue) + if _, err := persistence.QueueCommand(ctx, store.QueueCommandInput{IssueUUID: parsedIssue, ClientID: hello.GetClientId(), IssueTime: time.Now().UTC(), ReceiptTime: time.Now().UTC(), ImmutableSHA256: hash, ExecutionSpec: serialized}); err != nil { + t.Fatal(err) + } + if _, err := agent.Reconcile(ctx, transport, accepted, &rvboxv1.ReconcileSnapshot{}, agentproto.DefaultLimits()); err != nil { + t.Fatal(err) + } + readContext, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + encoded, err := transport.Read(readContext) + if err != nil { + t.Fatal(err) + } + dispatch, err := agentproto.DecodeEnvelope(encoded, agentproto.DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil || dispatch.GetCommandDispatch() == nil || dispatch.GetCommandDispatch().GetIssueUuid() != issue { + t.Fatalf("initial dispatch = %#v, %v", dispatch, err) + } + acceptedEnvelope, err := proto.Marshal(&rvboxv1.AgentEnvelope{SessionId: accepted.ID, SessionGeneration: accepted.Generation, Payload: &rvboxv1.AgentEnvelope_CommandAccepted{CommandAccepted: &rvboxv1.CommandAccepted{IssueUuid: issue, CommandRevision: 1, Accepted: true}}}) + if err != nil { + t.Fatal(err) + } + if err := transport.Write(ctx, acceptedEnvelope); err != nil { + t.Fatal(err) + } + stdinID := fixedIssueID(0x7b) + if _, err := controlService.AppendStdin(ctx, &rvboxv1.AppendStdinRequest{ClientId: hello.GetClientId(), IssueUuid: issue, RequestId: stdinID, Data: []byte("input"), AppendNewline: true}); err != nil { + t.Fatal(err) + } + if intents, pendingErr := persistence.PendingStdin(ctx, hello.GetClientId()); pendingErr != nil { + t.Fatalf("pending stdin query = %v", pendingErr) + } else if len(intents) != 1 || string(intents[0].Data) != "input" { + t.Fatalf("pending stdin intents = %#v", intents) + } + encoded, err = transport.Read(readContext) + if err != nil { + t.Fatal(err) + } + stdinEnvelope, err := agentproto.DecodeEnvelope(encoded, agentproto.DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil || stdinEnvelope.GetStdinWrite() == nil || stdinEnvelope.GetStdinWrite().GetWriteSeq() != 1 || string(stdinEnvelope.GetStdinWrite().GetData()) != "input" || !stdinEnvelope.GetStdinWrite().GetAppendNewline() { + t.Fatalf("stdin delivery = %#v, %v", stdinEnvelope, err) + } + ackEvent := &rvboxv1.CommandEvent{IssueUuid: issue, EventSeq: 1, ObservedAt: timestamppb.Now(), Payload: &rvboxv1.CommandEvent_StdinAck{StdinAck: &rvboxv1.StdinAcknowledgement{WriteSeq: 1}}} + if err := agent.SendCommandEvent(ctx, transport, accepted, ackEvent, agentproto.DefaultLimits()); err != nil { + t.Fatal(err) + } + if _, err := transport.Read(readContext); err != nil { + t.Fatal(err) + } + var acknowledged int + if err := persistence.DB().QueryRow(`SELECT acknowledged FROM stdin_writes WHERE issue_uuid = ? AND write_seq = 1`, parsedIssue[:]).Scan(&acknowledged); err != nil || acknowledged != 1 { + t.Fatalf("stdin acknowledgement = %d, %v", acknowledged, err) + } + closeID := fixedIssueID(0x7c) + if _, err := controlService.CloseStdin(ctx, &rvboxv1.CloseStdinRequest{ClientId: hello.GetClientId(), IssueUuid: issue, RequestId: closeID}); err != nil { + t.Fatal(err) + } + encoded, err = transport.Read(readContext) + if err != nil { + t.Fatal(err) + } + closeEnvelope, err := agentproto.DecodeEnvelope(encoded, agentproto.DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil || closeEnvelope.GetCloseStdin() == nil || closeEnvelope.GetCloseStdin().GetWriteSeq() != 2 { + t.Fatalf("stdin close delivery = %#v, %v", closeEnvelope, err) + } + closeAck := &rvboxv1.CommandEvent{IssueUuid: issue, EventSeq: 2, ObservedAt: timestamppb.Now(), Payload: &rvboxv1.CommandEvent_StdinAck{StdinAck: &rvboxv1.StdinAcknowledgement{WriteSeq: 2, StdinClosed: true}}} + if err := agent.SendCommandEvent(ctx, transport, accepted, closeAck, agentproto.DefaultLimits()); err != nil { + t.Fatal(err) + } + if _, err := transport.Read(readContext); err != nil { + t.Fatal(err) + } + var closeAcknowledged int + if err := persistence.DB().QueryRow(`SELECT acknowledged FROM stdin_writes WHERE issue_uuid = ? AND write_seq = 2`, parsedIssue[:]).Scan(&closeAcknowledged); err != nil || closeAcknowledged != 1 { + t.Fatalf("stdin close acknowledgement = %d, %v", closeAcknowledged, err) + } +} + func fixedIssueID(last byte) string { issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000001") issue[15] = last