diff --git a/internal/client/spool/migrations.go b/internal/client/spool/migrations.go index 7dbb2fe..bfffe86 100644 --- a/internal/client/spool/migrations.go +++ b/internal/client/spool/migrations.go @@ -84,6 +84,19 @@ CREATE TABLE events ( UNIQUE(issue_uuid, event_seq) ) STRICT, WITHOUT ROWID; 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 ( singleton INTEGER PRIMARY KEY CHECK(singleton = 1), client_total_charged_bytes INTEGER NOT NULL CHECK(client_total_charged_bytes >= 0), diff --git a/internal/client/spool/script.go b/internal/client/spool/script.go new file mode 100644 index 0000000..cd0e2f4 --- /dev/null +++ b/internal/client/spool/script.go @@ -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 +} diff --git a/internal/client/spool/script_test.go b/internal/client/spool/script_test.go new file mode 100644 index 0000000..31a9bd5 --- /dev/null +++ b/internal/client/spool/script_test.go @@ -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) + } +} diff --git a/internal/client/spool/spool.go b/internal/client/spool/spool.go index a90029c..11db4a7 100644 --- a/internal/client/spool/spool.go +++ b/internal/client/spool/spool.go @@ -38,6 +38,7 @@ type Options struct { BusyTimeout time.Duration TombstoneLimit uint64 QuotaLimits QuotaLimits + MaxScriptBytes uint64 } type Store struct { @@ -47,6 +48,7 @@ type Store struct { identity domain.UUID tombstoneLimit uint64 quotaLimits QuotaLimits + maxScriptBytes uint64 mu sync.Mutex } @@ -69,6 +71,15 @@ func Open(ctx context.Context, options Options) (*Store, error) { if err := options.QuotaLimits.Validate(); err != nil { 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 { return nil, err } @@ -100,7 +111,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} + 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 { _ = store.Close() return nil, fmt.Errorf("open client spool SQLite: %w", err) diff --git a/test/coverage.toml b/test/coverage.toml index af42676..d6ce4ad 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -195,7 +195,10 @@ tests = ["internal/agentproto/validate_test.go:TestCommandDispatchUUIDRevisionAn id = "BH-SCRIPT-01" layer = "unit" 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]] id = "BH-LAUNCH-01" @@ -224,6 +227,12 @@ layer = "unit" status = "implemented" 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]] id = "HP-WINCTX-02" layer = "integration"