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
+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
}