feat: persist and dispatch client command input

This commit is contained in:
2026-09-06 10:40:23 +00:00
parent 8155e5f81f
commit ad564de23f
16 changed files with 836 additions and 22 deletions
+6 -1
View File
@@ -9,6 +9,7 @@ import (
"github.com/rvbox/rvbox/internal/agentproto" "github.com/rvbox/rvbox/internal/agentproto"
"github.com/rvbox/rvbox/internal/client/spool" "github.com/rvbox/rvbox/internal/client/spool"
"github.com/rvbox/rvbox/internal/domain" "github.com/rvbox/rvbox/internal/domain"
"google.golang.org/protobuf/proto"
) )
// PersistDispatch admits a server dispatch to the durable spool before any // 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{} { if immutable == [32]byte{} {
return spool.Acceptance{}, errors.New("empty dispatch immutable hash") 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)
} }
+3 -1
View File
@@ -1,6 +1,7 @@
package agent package agent
import ( import (
"bytes"
"context" "context"
"errors" "errors"
"path/filepath" "path/filepath"
@@ -92,7 +93,8 @@ func TestPersistDispatchUsesImmutableHash_HP_DISPATCH_04(t *testing.T) {
defer store.Close() 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"}}} 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()) 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) t.Fatalf("PersistDispatch = %#v, %v", accepted, err)
} }
duplicate, err := PersistDispatch(context.Background(), store, Session{Generation: 7}, dispatch, time.Now(), agentproto.DefaultLimits()) duplicate, err := PersistDispatch(context.Background(), store, Session{Generation: 7}, dispatch, time.Now(), agentproto.DefaultLimits())
+99 -5
View File
@@ -1,12 +1,15 @@
package spool package spool
import ( import (
"bytes"
"context" "context"
"crypto/sha256"
"database/sql" "database/sql"
"errors" "errors"
"fmt" "fmt"
"time" "time"
"github.com/klauspost/compress/zstd"
"github.com/rvbox/rvbox/internal/domain" "github.com/rvbox/rvbox/internal/domain"
) )
@@ -16,6 +19,10 @@ type Command struct {
Revision uint64 Revision uint64
Phase uint32 Phase uint32
Terminal bool 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 { 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() { 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") 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) tx, err := store.db.BeginTx(ctx, nil)
if err != nil { if err != nil {
return Acceptance{}, err return Acceptance{}, err
} }
defer tx.Rollback() defer tx.Rollback()
existing, found, err := commandByUUID(ctx, tx, command.IssueUUID) existing, found, err := commandByUUID(ctx, tx, command.IssueUUID, store.maxExecutionSpecBytes)
if err != nil { if err != nil {
return Acceptance{}, err return Acceptance{}, err
} }
@@ -68,6 +91,9 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted
if existing.ImmutableSHA256 != command.ImmutableSHA256 { if existing.ImmutableSHA256 != command.ImmutableSHA256 {
return Acceptance{}, ErrCommandConflict return Acceptance{}, ErrCommandConflict
} }
if len(command.ExecutionSpec) > 0 && !bytes.Equal(existing.ExecutionSpec, command.ExecutionSpec) {
return Acceptance{}, ErrCommandConflict
}
return Acceptance{Duplicate: true, Command: existing}, nil return Acceptance{Duplicate: true, Command: existing}, nil
} }
var tombstoneHash []byte var tombstoneHash []byte
@@ -81,7 +107,7 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted
if err != sql.ErrNoRows { if err != sql.ErrNoRows {
return Acceptance{}, err return Acceptance{}, err
} }
charge, err := EstimateCharge(ChargeInput{SQLiteRows: 1, IndexEntries: 2}) baseCharge, err := EstimateCharge(ChargeInput{SQLiteRows: 1, IndexEntries: 2})
if err != nil { if err != nil {
return Acceptance{}, err return Acceptance{}, err
} }
@@ -89,23 +115,51 @@ func (store *Store) AcceptCommand(ctx context.Context, command Command, accepted
if err != nil { if err != nil {
return Acceptance{}, err 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}) decision, err := CheckReservation(store.quotaLimits, ReservationState{ClientTotalCharged: clientTotal, CloseoutRemaining: store.quotaLimits.CloseoutReserveBytes}, ReservationRequest{ChargedBytes: charge})
if err != nil { if err != nil {
return Acceptance{}, err 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 { if err != nil {
return Acceptance{}, err return Acceptance{}, err
} }
if err := updateClientTotalCharge(ctx, tx, decision.ClientTotalCharged); err != nil { if err := updateClientTotalCharge(ctx, tx, decision.ClientTotalCharged); err != nil {
return Acceptance{}, err 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 { if err := tx.Commit(); err != nil {
return Acceptance{}, err return Acceptance{}, err
} }
return Acceptance{Command: command}, nil 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 // AppendEvent assigns command-local order only. EventSeq remains zero until
// AssignSendWindow makes the durable event eligible for a wire send. // AssignSendWindow makes the durable event eligible for a wire send.
func (store *Store) AppendEvent(ctx context.Context, issueUUID domain.UUID, input EventInput) (Event, error) { 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 { func commandByUUID(ctx context.Context, query interface {
QueryRowContext(context.Context, string, ...any) *sql.Row 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 command Command
var hash []byte var hash []byte
var terminal int 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 { if err == sql.ErrNoRows {
return Command{}, false, nil return Command{}, false, nil
} }
@@ -465,9 +523,45 @@ func commandByUUID(ctx context.Context, query interface {
copy(command.ImmutableSHA256[:], hash) copy(command.ImmutableSHA256[:], hash)
command.IssueUUID = issueUUID command.IssueUUID = issueUUID
command.Terminal = terminal != 0 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 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) { func tombstonesToTrim(ctx context.Context, tx *sql.Tx, limit uint64) (uint64, error) {
var count uint64 var count uint64
if err := tx.QueryRowContext(ctx, `SELECT count(*) FROM command_tombstones`).Scan(&count); err != nil { if err := tx.QueryRowContext(ctx, `SELECT count(*) FROM command_tombstones`).Scan(&count); err != nil {
+20 -1
View File
@@ -13,7 +13,10 @@ type migration struct {
sql string 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 { func applyMigrations(ctx context.Context, db *sql.DB) error {
if _, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations ( if _, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (
@@ -113,3 +116,19 @@ CREATE TABLE command_tombstones (
) STRICT; ) STRICT;
CREATE INDEX tombstones_fifo ON command_tombstones(acknowledged_at, issue_uuid); 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;
`
+33 -2
View File
@@ -19,6 +19,7 @@ var (
type RecoveryReport struct { type RecoveryReport struct {
EventsChecked uint64 EventsChecked uint64
ScriptsChecked uint64 ScriptsChecked uint64
SpecsChecked uint64
} }
// Check verifies state that SQLite's own structural check cannot express. The // 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 { if err := checkQuotaCounters(ctx, store.db); err != nil {
return RecoveryReport{}, err 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 { 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 err := database.QueryRowContext(ctx, `SELECT count(*) FROM commands AS c
WHERE c.total_charged_bytes != c.base_charged_bytes + 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 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) COALESCE((SELECT charged_bytes FROM scripts WHERE issue_uuid = c.issue_uuid), 0)
OR c.output_charged_bytes != OR c.output_charged_bytes !=
COALESCE((SELECT sum(charged_bytes) FROM events WHERE issue_uuid = c.issue_uuid AND output = 1), 0)`).Scan(&invalid) 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 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 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`) 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 { if err != nil {
@@ -119,6 +121,35 @@ func checkPayloads(ctx context.Context, database *sql.DB, maxScriptBytes uint64)
if err := rows.Close(); err != nil { if err := rows.Close(); err != nil {
return report, err 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`) scripts, err := database.QueryContext(ctx, `SELECT issue_uuid, received_raw_bytes, stored_bytes, stored_data FROM scripts ORDER BY issue_uuid`)
if err != nil { if err != nil {
return report, err return report, err
+24 -10
View File
@@ -37,18 +37,23 @@ type Options struct {
BusyTimeout time.Duration BusyTimeout time.Duration
TombstoneLimit uint64 TombstoneLimit uint64
QuotaLimits QuotaLimits 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 { type Store struct {
db *sql.DB db *sql.DB
unlock func() error unlock func() error
dataDir string dataDir string
identity domain.UUID identity domain.UUID
tombstoneLimit uint64 tombstoneLimit uint64
quotaLimits QuotaLimits quotaLimits QuotaLimits
maxScriptBytes uint64 maxScriptBytes uint64
mu sync.Mutex maxExecutionSpecBytes uint64
mu sync.Mutex
} }
// Open recovers an existing spool or creates an empty one. The caller must // 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 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 { if options.MaxScriptBytes > options.QuotaLimits.CommandTotalBytes {
return nil, errors.New("script maximum exceeds client command quota") 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.SetMaxOpenConns(1)
db.SetMaxIdleConns(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 { if err := db.PingContext(ctx); err != nil {
_ = store.Close() _ = store.Close()
return nil, fmt.Errorf("open client spool SQLite: %w", err) return nil, fmt.Errorf("open client spool SQLite: %w", err)
+64
View File
@@ -1,6 +1,7 @@
package spool package spool
import ( import (
"bytes"
"context" "context"
"crypto/sha256" "crypto/sha256"
"errors" "errors"
@@ -12,6 +13,69 @@ import (
"github.com/rvbox/rvbox/internal/domain" "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) { func TestStoreDurableIdentityAndSingleInstance_HP_CLIENT_01(t *testing.T) {
t.Parallel() t.Parallel()
ctx := context.Background() ctx := context.Background()
+101
View File
@@ -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
}
@@ -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)
}
}
@@ -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, '=')
}
@@ -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)
}
}
+66 -2
View File
@@ -141,13 +141,14 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w
capacity := &CapacityShadow{} capacity := &CapacityShadow{}
var capacityMu sync.Mutex var capacityMu sync.Mutex
reservations := make(map[string]DispatchLane) reservations := make(map[string]DispatchLane)
stdinSent := make(map[string]struct{})
reconciled := make(chan struct{}) reconciled := make(chan struct{})
var reconcileOnce sync.Once var reconcileOnce sync.Once
if !capacity.UpdateAdvertised(0, 0, hello.GetMaxRunningCommands(), hello.GetMaxQueuedCommands()) { if !capacity.UpdateAdvertised(0, 0, hello.GetMaxRunningCommands(), hello.GetMaxQueuedCommands()) {
server.close(connection, websocket.StatusPolicyViolation, "invalid initial client capacity") server.close(connection, websocket.StatusPolicyViolation, "invalid initial client capacity")
return 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) encodedSessionID := encodeSessionID(sessionID)
welcome, err := proto.Marshal(&rvboxv1.AgentEnvelope{ welcome, err := proto.Marshal(&rvboxv1.AgentEnvelope{
SessionId: encodedSessionID, SessionGeneration: registration.Generation, 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") server.close(connection, websocket.StatusInternalError, "could not acknowledge command event")
return 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 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 // dispatchLoop is the per-session serialized dispatcher. It waits for a
// complete reconciliation result before consuming queued work, then coalesces // complete reconciliation result before consuming queued work, then coalesces
// wakeups from local control RPCs, capacity advertisements, and acceptances. // 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 { select {
case <-reconciled: case <-reconciled:
case <-ctx.Done(): case <-ctx.Done():
@@ -377,6 +433,14 @@ func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue,
} }
for { for {
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() capacityMu.Lock()
lane := capacity.Reserve() lane := capacity.Reserve()
capacityMu.Unlock() capacityMu.Unlock()
+13
View File
@@ -12,6 +12,7 @@ import (
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
"github.com/rvbox/rvbox/internal/domain" "github.com/rvbox/rvbox/internal/domain"
"google.golang.org/protobuf/proto"
) )
const ( 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 { if err == nil {
var update sql.Result var update sql.Result
update, err = tx.ExecContext(ctx, `UPDATE clients SET charged_bytes = ? WHERE client_id = ? AND charged_bytes = ?`, reservation.ClientTotalCharged, clientID, clientCharged) update, err = tx.ExecContext(ctx, `UPDATE clients SET charged_bytes = ? WHERE client_id = ? AND charged_bytes = ?`, reservation.ClientTotalCharged, clientID, clientCharged)
+80
View File
@@ -1,6 +1,7 @@
package store package store
import ( import (
"bytes"
"context" "context"
"crypto/sha256" "crypto/sha256"
"database/sql" "database/sql"
@@ -9,6 +10,7 @@ import (
"math" "math"
"time" "time"
"github.com/klauspost/compress/zstd"
"github.com/rvbox/rvbox/internal/domain" "github.com/rvbox/rvbox/internal/domain"
) )
@@ -33,6 +35,84 @@ type StdinWriteResult struct {
Duplicate bool 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 // AppendStdin durably records one ordered stdin intent. It does not claim
// delivery to a process; an agent acknowledgement is a later protocol event. // delivery to a process; an agent acknowledgement is a later protocol event.
// Reusing RequestUUID with the same immutable hash returns the original write // Reusing RequestUUID with the same immutable hash returns the original write
+36
View File
@@ -227,6 +227,24 @@ layer = "unit"
status = "implemented" status = "implemented"
tests = ["internal/client/supervisor/windows/selection_test.go:TestSelectExecutionContextBoundaries_BH_WINCTX_02"] 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]] [[requirements]]
id = "BH-SES-01" id = "BH-SES-01"
layer = "unit" layer = "unit"
@@ -302,6 +320,24 @@ layer = "integration"
status = "implemented" status = "implemented"
tests = ["test/integration/clientagent/clientagent_integration_test.go:TestControlQueueWakesReconciledSession_HP_DISPATCH_06"] 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]] [[requirements]]
id = "HP-EVENT-01" id = "HP-EVENT-01"
layer = "unit" layer = "unit"
@@ -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 { func fixedIssueID(last byte) string {
issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000001") issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000001")
issue[15] = last issue[15] = last