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 }