feat: persist resumable client script uploads

This commit is contained in:
2026-09-06 06:39:12 +00:00
parent 6539b9e965
commit 62c0052979
5 changed files with 447 additions and 2 deletions
+13
View File
@@ -84,6 +84,19 @@ CREATE TABLE events (
UNIQUE(issue_uuid, event_seq) UNIQUE(issue_uuid, event_seq)
) STRICT, WITHOUT ROWID; ) STRICT, WITHOUT ROWID;
CREATE INDEX events_send_window ON events(issue_uuid, event_seq, local_ordinal); CREATE INDEX events_send_window ON events(issue_uuid, event_seq, local_ordinal);
CREATE TABLE scripts (
issue_uuid BLOB PRIMARY KEY REFERENCES commands(issue_uuid) ON DELETE CASCADE CHECK(length(issue_uuid) = 16),
declared_raw_bytes INTEGER NOT NULL CHECK(declared_raw_bytes >= 0),
declared_sha256 BLOB NOT NULL CHECK(length(declared_sha256) = 32),
received_raw_bytes INTEGER NOT NULL DEFAULT 0 CHECK(received_raw_bytes >= 0),
stored_bytes INTEGER NOT NULL CHECK(stored_bytes >= 0),
compression INTEGER NOT NULL CHECK(compression = 2),
stored_data BLOB NOT NULL,
charged_bytes INTEGER NOT NULL CHECK(charged_bytes > 0),
committed INTEGER NOT NULL DEFAULT 0 CHECK(committed IN (0, 1)),
CHECK(received_raw_bytes <= declared_raw_bytes),
CHECK((committed = 0) OR received_raw_bytes = declared_raw_bytes)
) STRICT, WITHOUT ROWID;
CREATE TABLE spool_counters ( CREATE TABLE spool_counters (
singleton INTEGER PRIMARY KEY CHECK(singleton = 1), singleton INTEGER PRIMARY KEY CHECK(singleton = 1),
client_total_charged_bytes INTEGER NOT NULL CHECK(client_total_charged_bytes >= 0), client_total_charged_bytes INTEGER NOT NULL CHECK(client_total_charged_bytes >= 0),
+302
View File
@@ -0,0 +1,302 @@
package spool
import (
"bytes"
"context"
"crypto/sha256"
"database/sql"
"errors"
"io"
"github.com/klauspost/compress/zstd"
"github.com/rvbox/rvbox/internal/domain"
)
var (
ErrScriptConflict = errors.New("script descriptor conflicts with accepted command")
ErrScriptHole = errors.New("script chunk has a gap")
ErrScriptOverlap = errors.New("script chunk is not an exact replay")
ErrScriptDigest = errors.New("script digest does not match")
ErrScriptBounds = errors.New("script chunk is outside declared bounds")
ErrScriptState = errors.New("stored script state is corrupt")
ErrScriptTerminal = errors.New("terminal command cannot accept script data")
)
type ScriptDescriptor struct {
SizeBytes uint64
SHA256 [sha256.Size]byte
}
type ScriptStatus struct {
ReceivedBytes uint64
Committed bool
Duplicate bool
}
// BeginScript persists the immutable descriptor at command acceptance time. A
// matching replay is harmless; a different descriptor is a protocol conflict.
func (store *Store) BeginScript(ctx context.Context, issueUUID domain.UUID, descriptor ScriptDescriptor) (ScriptStatus, error) {
if !validUUID(issueUUID) || descriptor.SizeBytes > store.maxScriptBytes {
return ScriptStatus{}, ErrScriptBounds
}
tx, err := store.db.BeginTx(ctx, nil)
if err != nil {
return ScriptStatus{}, err
}
defer tx.Rollback()
if err := commandMayAcceptScript(ctx, tx, issueUUID); err != nil {
return ScriptStatus{}, err
}
status, found, err := scriptStatus(ctx, tx, issueUUID)
if err != nil {
return ScriptStatus{}, err
}
if found {
var size uint64
var digest []byte
if err := tx.QueryRowContext(ctx, `SELECT declared_raw_bytes, declared_sha256 FROM scripts WHERE issue_uuid = ?`, issueUUID[:]).Scan(&size, &digest); err != nil {
return ScriptStatus{}, err
}
if size != descriptor.SizeBytes || !bytes.Equal(digest, descriptor.SHA256[:]) {
return ScriptStatus{}, ErrScriptConflict
}
status.Duplicate = true
return status, tx.Commit()
}
stored, err := compressScript(nil)
if err != nil {
return ScriptStatus{}, err
}
charge, err := EstimateCharge(ChargeInput{EncodedBytes: uint64(len(stored)), SQLiteRows: 1, IndexEntries: 1})
if err != nil {
return ScriptStatus{}, err
}
if err := store.applyNonOutputCharge(ctx, tx, issueUUID, 0, charge); err != nil {
return ScriptStatus{}, err
}
if _, err := tx.ExecContext(ctx, `INSERT INTO scripts(issue_uuid, declared_raw_bytes, declared_sha256, stored_bytes, compression, stored_data, charged_bytes) VALUES (?, ?, ?, ?, 2, ?, ?)`, issueUUID[:], descriptor.SizeBytes, descriptor.SHA256[:], len(stored), stored, charge); err != nil {
return ScriptStatus{}, err
}
if err := tx.Commit(); err != nil {
return ScriptStatus{}, err
}
return ScriptStatus{}, nil
}
// AppendScriptChunk accepts only the next contiguous bytes or a complete exact
// replay of already durable bytes. It never fills a gap or accepts an overlap
// that extends the durable prefix.
func (store *Store) AppendScriptChunk(ctx context.Context, issueUUID domain.UUID, offset uint64, data []byte, digest [sha256.Size]byte) (ScriptStatus, error) {
if !validUUID(issueUUID) || len(data) == 0 || sha256.Sum256(data) != digest {
return ScriptStatus{}, ErrScriptDigest
}
tx, err := store.db.BeginTx(ctx, nil)
if err != nil {
return ScriptStatus{}, err
}
defer tx.Rollback()
if err := commandMayAcceptScript(ctx, tx, issueUUID); err != nil {
return ScriptStatus{}, err
}
row, err := loadScript(ctx, tx, issueUUID)
if err != nil {
return ScriptStatus{}, err
}
if row.Committed {
return ScriptStatus{}, ErrScriptTerminal
}
if offset > row.ReceivedBytes {
return ScriptStatus{}, ErrScriptHole
}
if offset > row.DeclaredBytes {
return ScriptStatus{}, ErrScriptBounds
}
if uint64(len(data)) > row.DeclaredBytes-offset {
return ScriptStatus{}, ErrScriptBounds
}
raw, err := decompressScript(row.Stored, row.ReceivedBytes, store.maxScriptBytes)
if err != nil {
return ScriptStatus{}, err
}
end := offset + uint64(len(data))
if offset < row.ReceivedBytes {
if end > row.ReceivedBytes || !bytes.Equal(raw[offset:end], data) {
return ScriptStatus{}, ErrScriptOverlap
}
return ScriptStatus{ReceivedBytes: row.ReceivedBytes, Duplicate: true}, tx.Commit()
}
updatedRaw := append(raw, data...)
stored, err := compressScript(updatedRaw)
if err != nil {
return ScriptStatus{}, err
}
charge, err := EstimateCharge(ChargeInput{EncodedBytes: uint64(len(stored)), SQLiteRows: 1, IndexEntries: 1})
if err != nil {
return ScriptStatus{}, err
}
if err := store.applyNonOutputCharge(ctx, tx, issueUUID, row.ChargedBytes, charge); err != nil {
return ScriptStatus{}, err
}
if _, err := tx.ExecContext(ctx, `UPDATE scripts SET received_raw_bytes = ?, stored_bytes = ?, stored_data = ?, charged_bytes = ? WHERE issue_uuid = ?`, uint64(len(updatedRaw)), len(stored), stored, charge, issueUUID[:]); err != nil {
return ScriptStatus{}, err
}
if err := tx.Commit(); err != nil {
return ScriptStatus{}, err
}
return ScriptStatus{ReceivedBytes: uint64(len(updatedRaw))}, nil
}
// CommitScript is durable launch eligibility, not process authorization. The
// supervisor must still persist launch_prepared/launch_authorized separately.
func (store *Store) CommitScript(ctx context.Context, issueUUID domain.UUID, descriptor ScriptDescriptor) (ScriptStatus, error) {
tx, err := store.db.BeginTx(ctx, nil)
if err != nil {
return ScriptStatus{}, err
}
defer tx.Rollback()
if err := commandMayAcceptScript(ctx, tx, issueUUID); err != nil {
return ScriptStatus{}, err
}
row, err := loadScript(ctx, tx, issueUUID)
if err != nil {
return ScriptStatus{}, err
}
if row.DeclaredBytes != descriptor.SizeBytes || row.DeclaredSHA256 != descriptor.SHA256 {
return ScriptStatus{}, ErrScriptConflict
}
if row.Committed {
return ScriptStatus{ReceivedBytes: row.ReceivedBytes, Committed: true, Duplicate: true}, tx.Commit()
}
if row.ReceivedBytes != row.DeclaredBytes {
return ScriptStatus{ReceivedBytes: row.ReceivedBytes}, ErrScriptBounds
}
raw, err := decompressScript(row.Stored, row.ReceivedBytes, store.maxScriptBytes)
if err != nil {
return ScriptStatus{}, err
}
if sha256.Sum256(raw) != descriptor.SHA256 {
return ScriptStatus{}, ErrScriptDigest
}
if _, err := tx.ExecContext(ctx, `UPDATE scripts SET committed = 1 WHERE issue_uuid = ?`, issueUUID[:]); err != nil {
return ScriptStatus{}, err
}
if err := tx.Commit(); err != nil {
return ScriptStatus{}, err
}
return ScriptStatus{ReceivedBytes: row.ReceivedBytes, Committed: true}, nil
}
type storedScript struct {
DeclaredBytes uint64
DeclaredSHA256 [sha256.Size]byte
ReceivedBytes uint64
Stored []byte
ChargedBytes uint64
Committed bool
}
func loadScript(ctx context.Context, tx *sql.Tx, issueUUID domain.UUID) (storedScript, error) {
var result storedScript
var digest []byte
var storedBytes uint64
var compression uint32
var committed int
err := tx.QueryRowContext(ctx, `SELECT declared_raw_bytes, declared_sha256, received_raw_bytes, stored_bytes, compression, stored_data, charged_bytes, committed FROM scripts WHERE issue_uuid = ?`, issueUUID[:]).Scan(&result.DeclaredBytes, &digest, &result.ReceivedBytes, &storedBytes, &compression, &result.Stored, &result.ChargedBytes, &committed)
if err == sql.ErrNoRows {
return result, ErrScriptBounds
}
if err != nil {
return result, err
}
if len(digest) != sha256.Size || storedBytes != uint64(len(result.Stored)) || compression != 2 || result.ReceivedBytes > result.DeclaredBytes {
return storedScript{}, ErrScriptState
}
copy(result.DeclaredSHA256[:], digest)
result.Committed = committed != 0
return result, nil
}
func scriptStatus(ctx context.Context, tx *sql.Tx, issueUUID domain.UUID) (ScriptStatus, bool, error) {
var status ScriptStatus
var committed int
err := tx.QueryRowContext(ctx, `SELECT received_raw_bytes, committed FROM scripts WHERE issue_uuid = ?`, issueUUID[:]).Scan(&status.ReceivedBytes, &committed)
if err == sql.ErrNoRows {
return ScriptStatus{}, false, nil
}
if err != nil {
return ScriptStatus{}, false, err
}
status.Committed = committed != 0
return status, true, nil
}
func commandMayAcceptScript(ctx context.Context, tx *sql.Tx, issueUUID domain.UUID) error {
var terminal int
err := tx.QueryRowContext(ctx, `SELECT terminal FROM commands WHERE issue_uuid = ?`, issueUUID[:]).Scan(&terminal)
if err == sql.ErrNoRows {
return ErrUnknownCommand
}
if err != nil {
return err
}
if terminal != 0 {
return ErrScriptTerminal
}
return nil
}
func (store *Store) applyNonOutputCharge(ctx context.Context, tx *sql.Tx, issueUUID domain.UUID, oldCharge, newCharge uint64) error {
var outputCharged, totalCharged, closeout uint64
if err := tx.QueryRowContext(ctx, `SELECT output_charged_bytes, total_charged_bytes, closeout_remaining_bytes FROM commands WHERE issue_uuid = ?`, issueUUID[:]).Scan(&outputCharged, &totalCharged, &closeout); err != nil {
if err == sql.ErrNoRows {
return ErrUnknownCommand
}
return err
}
clientTotal, err := clientTotalCharge(ctx, tx)
if err != nil {
return err
}
if oldCharge > totalCharged || oldCharge > clientTotal {
return ErrScriptState
}
if newCharge >= oldCharge {
if newCharge == oldCharge {
return nil
}
decision, err := CheckReservation(store.quotaLimits, ReservationState{CommandOutputCharged: outputCharged, CommandTotalCharged: totalCharged, ClientTotalCharged: clientTotal, CloseoutRemaining: closeout}, ReservationRequest{ChargedBytes: newCharge - oldCharge})
if err != nil {
return err
}
if _, err := tx.ExecContext(ctx, `UPDATE commands SET total_charged_bytes = ?, closeout_remaining_bytes = ? WHERE issue_uuid = ?`, decision.CommandTotalCharged, decision.CloseoutRemaining, issueUUID[:]); err != nil {
return err
}
return updateClientTotalCharge(ctx, tx, decision.ClientTotalCharged)
}
if _, err := tx.ExecContext(ctx, `UPDATE commands SET total_charged_bytes = ? WHERE issue_uuid = ?`, totalCharged-(oldCharge-newCharge), issueUUID[:]); err != nil {
return err
}
return updateClientTotalCharge(ctx, tx, clientTotal-(oldCharge-newCharge))
}
func compressScript(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 decompressScript(stored []byte, expected, maximum uint64) ([]byte, error) {
decoder, err := zstd.NewReader(bytes.NewReader(stored), zstd.WithDecoderConcurrency(1), zstd.WithDecoderMaxMemory(maximum+1))
if err != nil {
return nil, ErrScriptState
}
defer decoder.Close()
raw, err := io.ReadAll(io.LimitReader(decoder, int64(maximum)+1))
if err != nil || uint64(len(raw)) != expected || uint64(len(raw)) > maximum {
return nil, ErrScriptState
}
return raw, nil
}
+110
View File
@@ -0,0 +1,110 @@
package spool
import (
"context"
"crypto/sha256"
"errors"
"path/filepath"
"testing"
"time"
)
func TestScriptUploadDurableExactReplayAndCommit_HP_SCRIPT_01(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-000000000031")
if _, err := store.AcceptCommand(ctx, testCommand(issue, []byte("script command")), time.Now().UTC()); err != nil {
t.Fatal(err)
}
raw := []byte("first chunk; second chunk")
descriptor := ScriptDescriptor{SizeBytes: uint64(len(raw)), SHA256: sha256.Sum256(raw)}
status, err := store.BeginScript(ctx, issue, descriptor)
if err != nil || status.ReceivedBytes != 0 || status.Committed {
t.Fatalf("BeginScript = %#v, %v", status, err)
}
duplicate, err := store.BeginScript(ctx, issue, descriptor)
if err != nil || !duplicate.Duplicate {
t.Fatalf("matching BeginScript replay = %#v, %v", duplicate, err)
}
first := raw[:12]
status, err = store.AppendScriptChunk(ctx, issue, 0, first, sha256.Sum256(first))
if err != nil || status.ReceivedBytes != uint64(len(first)) {
t.Fatalf("first ScriptChunk = %#v, %v", status, err)
}
duplicate, err = store.AppendScriptChunk(ctx, issue, 0, first, sha256.Sum256(first))
if err != nil || !duplicate.Duplicate || duplicate.ReceivedBytes != uint64(len(first)) {
t.Fatalf("matching ScriptChunk replay = %#v, %v", duplicate, err)
}
if _, err := store.AppendScriptChunk(ctx, issue, uint64(len(first)+1), raw[len(first):], sha256.Sum256(raw[len(first):])); !errors.Is(err, ErrScriptHole) {
t.Fatalf("script hole error = %v, want ErrScriptHole", err)
}
if _, err := store.AppendScriptChunk(ctx, issue, 0, []byte("different one"), sha256.Sum256([]byte("different one"))); !errors.Is(err, ErrScriptOverlap) {
t.Fatalf("changed overlap error = %v, want ErrScriptOverlap", err)
}
if _, err := store.CommitScript(ctx, issue, descriptor); !errors.Is(err, ErrScriptBounds) {
t.Fatalf("early commit error = %v, want ErrScriptBounds", err)
}
if err := store.Close(); err != nil {
t.Fatal(err)
}
store = openTestStore(t, ctx, directory, DefaultTombstoneLimit)
rest := raw[len(first):]
status, err = store.AppendScriptChunk(ctx, issue, uint64(len(first)), rest, sha256.Sum256(rest))
if err != nil || status.ReceivedBytes != uint64(len(raw)) {
t.Fatalf("resumed ScriptChunk = %#v, %v", status, err)
}
status, err = store.CommitScript(ctx, issue, descriptor)
if err != nil || !status.Committed {
t.Fatalf("ScriptCommit = %#v, %v", status, err)
}
duplicate, err = store.CommitScript(ctx, issue, descriptor)
if err != nil || !duplicate.Duplicate || !duplicate.Committed {
t.Fatalf("ScriptCommit replay = %#v, %v", duplicate, err)
}
if _, err := store.AppendScriptChunk(ctx, issue, uint64(len(raw)), []byte("!"), sha256.Sum256([]byte("!"))); !errors.Is(err, ErrScriptTerminal) {
t.Fatalf("post-commit chunk error = %v, want ErrScriptTerminal", err)
}
var compression uint32
var stored []byte
if err := store.db.QueryRowContext(ctx, `SELECT compression, stored_data FROM scripts WHERE issue_uuid = ?`, issue[:]).Scan(&compression, &stored); err != nil {
t.Fatal(err)
}
if compression != 2 || string(stored) == string(raw) {
t.Fatalf("script was not stored as zstd data: compression=%d stored=%q", compression, stored)
}
}
func TestScriptUploadRejectsDescriptorDigestAndBounds_BH_SCRIPT_01(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-000000000032")
if _, err := store.AcceptCommand(ctx, testCommand(issue, []byte("script command")), time.Now().UTC()); err != nil {
t.Fatal(err)
}
raw := []byte("abc")
descriptor := ScriptDescriptor{SizeBytes: uint64(len(raw)), SHA256: sha256.Sum256(raw)}
if _, err := store.BeginScript(ctx, issue, descriptor); err != nil {
t.Fatal(err)
}
wrong := sha256.Sum256([]byte("wrong"))
if _, err := store.AppendScriptChunk(ctx, issue, 0, raw, wrong); !errors.Is(err, ErrScriptDigest) {
t.Fatalf("chunk digest error = %v, want ErrScriptDigest", err)
}
if _, err := store.AppendScriptChunk(ctx, issue, 0, []byte("abcd"), sha256.Sum256([]byte("abcd"))); !errors.Is(err, ErrScriptBounds) {
t.Fatalf("out-of-bounds chunk error = %v, want ErrScriptBounds", err)
}
conflicting := descriptor
conflicting.SHA256 = wrong
if _, err := store.BeginScript(ctx, issue, conflicting); !errors.Is(err, ErrScriptConflict) {
t.Fatalf("descriptor conflict error = %v, want ErrScriptConflict", err)
}
if err := store.MarkTerminal(ctx, issue, 5); err != nil {
t.Fatal(err)
}
if _, err := store.AppendScriptChunk(ctx, issue, 0, raw, sha256.Sum256(raw)); !errors.Is(err, ErrScriptTerminal) {
t.Fatalf("terminal command chunk error = %v, want ErrScriptTerminal", err)
}
}
+12 -1
View File
@@ -38,6 +38,7 @@ type Options struct {
BusyTimeout time.Duration BusyTimeout time.Duration
TombstoneLimit uint64 TombstoneLimit uint64
QuotaLimits QuotaLimits QuotaLimits QuotaLimits
MaxScriptBytes uint64
} }
type Store struct { type Store struct {
@@ -47,6 +48,7 @@ type Store struct {
identity domain.UUID identity domain.UUID
tombstoneLimit uint64 tombstoneLimit uint64
quotaLimits QuotaLimits quotaLimits QuotaLimits
maxScriptBytes uint64
mu sync.Mutex mu sync.Mutex
} }
@@ -69,6 +71,15 @@ func Open(ctx context.Context, options Options) (*Store, error) {
if err := options.QuotaLimits.Validate(); err != nil { if err := options.QuotaLimits.Validate(); err != nil {
return nil, err return nil, err
} }
if options.MaxScriptBytes == 0 {
options.MaxScriptBytes = 10 << 20
if options.MaxScriptBytes > options.QuotaLimits.CommandTotalBytes {
options.MaxScriptBytes = options.QuotaLimits.CommandTotalBytes
}
}
if options.MaxScriptBytes > options.QuotaLimits.CommandTotalBytes {
return nil, errors.New("script maximum exceeds client command quota")
}
if err := ensurePrivateDirectory(options.DataDir); err != nil { if err := ensurePrivateDirectory(options.DataDir); err != nil {
return nil, err return nil, err
} }
@@ -100,7 +111,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} store := &Store{db: db, unlock: unlock, dataDir: options.DataDir, identity: identity, tombstoneLimit: options.TombstoneLimit, quotaLimits: options.QuotaLimits, maxScriptBytes: options.MaxScriptBytes}
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)
+10 -1
View File
@@ -195,7 +195,10 @@ tests = ["internal/agentproto/validate_test.go:TestCommandDispatchUUIDRevisionAn
id = "BH-SCRIPT-01" id = "BH-SCRIPT-01"
layer = "unit" layer = "unit"
status = "implemented" status = "implemented"
tests = ["internal/agentproto/validate_test.go:TestScriptDescriptorHostileNames_BH_SCRIPT_01"] tests = [
"internal/agentproto/validate_test.go:TestScriptDescriptorHostileNames_BH_SCRIPT_01",
"internal/client/spool/script_test.go:TestScriptUploadRejectsDescriptorDigestAndBounds_BH_SCRIPT_01",
]
[[requirements]] [[requirements]]
id = "BH-LAUNCH-01" id = "BH-LAUNCH-01"
@@ -224,6 +227,12 @@ layer = "unit"
status = "implemented" status = "implemented"
tests = ["internal/client/spool/quota_test.go:TestSpoolAcknowledgementReleasesQuota_HP_CLIENT_08"] tests = ["internal/client/spool/quota_test.go:TestSpoolAcknowledgementReleasesQuota_HP_CLIENT_08"]
[[requirements]]
id = "HP-SCRIPT-01"
layer = "unit"
status = "implemented"
tests = ["internal/client/spool/script_test.go:TestScriptUploadDurableExactReplayAndCommit_HP_SCRIPT_01"]
[[requirements]] [[requirements]]
id = "HP-WINCTX-02" id = "HP-WINCTX-02"
layer = "integration" layer = "integration"