303 lines
11 KiB
Go
303 lines
11 KiB
Go
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
|
|
}
|