mirror of
https://github.com/tendermint/tendermint.git
synced 2026-09-24 08:54:51 +00:00
test: add end-to-end testing framework (#5435)
Partial fix for #5291. For details, see [README.md](https://github.com/tendermint/tendermint/blob/erik/e2e-tests/test/e2e/README.md) and [RFC-001](https://github.com/tendermint/tendermint/blob/master/docs/rfc/rfc-001-end-to-end-testing.md). This only includes a single test case under `test/e2e/tests/`, as a proof of concept - additional test cases will be submitted separately. A randomized testnet generator will also be submitted separately, there a currently just a handful of static testnets under `test/e2e/networks/`. This will eventually replace the current P2P tests and run in CI.
This commit is contained in:
committed by
Erik Grinaker
parent
1b733ea28d
commit
a58454e788
@@ -0,0 +1,217 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/tendermint/tendermint/abci/example/code"
|
||||
abci "github.com/tendermint/tendermint/abci/types"
|
||||
"github.com/tendermint/tendermint/libs/log"
|
||||
"github.com/tendermint/tendermint/version"
|
||||
)
|
||||
|
||||
// Application is an ABCI application for use by end-to-end tests. It is a
|
||||
// simple key/value store for strings, storing data in memory and persisting
|
||||
// to disk as JSON, taking state sync snapshots if requested.
|
||||
type Application struct {
|
||||
abci.BaseApplication
|
||||
logger log.Logger
|
||||
state *State
|
||||
snapshots *SnapshotStore
|
||||
cfg *Config
|
||||
restoreSnapshot *abci.Snapshot
|
||||
restoreChunks [][]byte
|
||||
}
|
||||
|
||||
// NewApplication creates the application.
|
||||
func NewApplication(cfg *Config) (*Application, error) {
|
||||
state, err := NewState(filepath.Join(cfg.Dir, "state.json"), cfg.PersistInterval)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
snapshots, err := NewSnapshotStore(filepath.Join(cfg.Dir, "snapshots"))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Application{
|
||||
logger: log.NewTMLogger(log.NewSyncWriter(os.Stdout)),
|
||||
state: state,
|
||||
snapshots: snapshots,
|
||||
cfg: cfg,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Info implements ABCI.
|
||||
func (app *Application) Info(req abci.RequestInfo) abci.ResponseInfo {
|
||||
return abci.ResponseInfo{
|
||||
Version: version.ABCIVersion,
|
||||
AppVersion: 1,
|
||||
LastBlockHeight: int64(app.state.Height),
|
||||
LastBlockAppHash: app.state.Hash,
|
||||
}
|
||||
}
|
||||
|
||||
// Info implements ABCI.
|
||||
func (app *Application) InitChain(req abci.RequestInitChain) abci.ResponseInitChain {
|
||||
var err error
|
||||
app.state.initialHeight = uint64(req.InitialHeight)
|
||||
if len(req.AppStateBytes) > 0 {
|
||||
err = app.state.Import(0, req.AppStateBytes)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
resp := abci.ResponseInitChain{
|
||||
AppHash: app.state.Hash,
|
||||
}
|
||||
if resp.Validators, err = app.validatorUpdates(0); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// CheckTx implements ABCI.
|
||||
func (app *Application) CheckTx(req abci.RequestCheckTx) abci.ResponseCheckTx {
|
||||
_, _, err := parseTx(req.Tx)
|
||||
if err != nil {
|
||||
return abci.ResponseCheckTx{
|
||||
Code: code.CodeTypeEncodingError,
|
||||
Log: err.Error(),
|
||||
}
|
||||
}
|
||||
return abci.ResponseCheckTx{Code: code.CodeTypeOK, GasWanted: 1}
|
||||
}
|
||||
|
||||
// DeliverTx implements ABCI.
|
||||
func (app *Application) DeliverTx(req abci.RequestDeliverTx) abci.ResponseDeliverTx {
|
||||
key, value, err := parseTx(req.Tx)
|
||||
if err != nil {
|
||||
panic(err) // shouldn't happen since we verified it in CheckTx
|
||||
}
|
||||
app.state.Set(key, value)
|
||||
return abci.ResponseDeliverTx{Code: code.CodeTypeOK}
|
||||
}
|
||||
|
||||
// EndBlock implements ABCI.
|
||||
func (app *Application) EndBlock(req abci.RequestEndBlock) abci.ResponseEndBlock {
|
||||
var err error
|
||||
resp := abci.ResponseEndBlock{}
|
||||
if resp.ValidatorUpdates, err = app.validatorUpdates(uint64(req.Height)); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// Commit implements ABCI.
|
||||
func (app *Application) Commit() abci.ResponseCommit {
|
||||
height, hash, err := app.state.Commit()
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
if app.cfg.SnapshotInterval > 0 && height%app.cfg.SnapshotInterval == 0 {
|
||||
snapshot, err := app.snapshots.Create(app.state)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
logger.Info("Created state sync snapshot", "height", snapshot.Height)
|
||||
}
|
||||
retainHeight := int64(0)
|
||||
if app.cfg.RetainBlocks > 0 {
|
||||
retainHeight = int64(height - app.cfg.RetainBlocks + 1)
|
||||
}
|
||||
return abci.ResponseCommit{
|
||||
Data: hash,
|
||||
RetainHeight: retainHeight,
|
||||
}
|
||||
}
|
||||
|
||||
// Query implements ABCI.
|
||||
func (app *Application) Query(req abci.RequestQuery) abci.ResponseQuery {
|
||||
return abci.ResponseQuery{
|
||||
Height: int64(app.state.Height),
|
||||
Key: req.Data,
|
||||
Value: []byte(app.state.Get(string(req.Data))),
|
||||
}
|
||||
}
|
||||
|
||||
// ListSnapshots implements ABCI.
|
||||
func (app *Application) ListSnapshots(req abci.RequestListSnapshots) abci.ResponseListSnapshots {
|
||||
snapshots, err := app.snapshots.List()
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return abci.ResponseListSnapshots{Snapshots: snapshots}
|
||||
}
|
||||
|
||||
// LoadSnapshotChunk implements ABCI.
|
||||
func (app *Application) LoadSnapshotChunk(req abci.RequestLoadSnapshotChunk) abci.ResponseLoadSnapshotChunk {
|
||||
chunk, err := app.snapshots.LoadChunk(req.Height, req.Format, req.Chunk)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return abci.ResponseLoadSnapshotChunk{Chunk: chunk}
|
||||
}
|
||||
|
||||
// OfferSnapshot implements ABCI.
|
||||
func (app *Application) OfferSnapshot(req abci.RequestOfferSnapshot) abci.ResponseOfferSnapshot {
|
||||
if app.restoreSnapshot != nil {
|
||||
panic("A snapshot is already being restored")
|
||||
}
|
||||
app.restoreSnapshot = req.Snapshot
|
||||
app.restoreChunks = [][]byte{}
|
||||
return abci.ResponseOfferSnapshot{Result: abci.ResponseOfferSnapshot_ACCEPT}
|
||||
}
|
||||
|
||||
// ApplySnapshotChunk implements ABCI.
|
||||
func (app *Application) ApplySnapshotChunk(req abci.RequestApplySnapshotChunk) abci.ResponseApplySnapshotChunk {
|
||||
if app.restoreSnapshot == nil {
|
||||
panic("No restore in progress")
|
||||
}
|
||||
app.restoreChunks = append(app.restoreChunks, req.Chunk)
|
||||
if len(app.restoreChunks) == int(app.restoreSnapshot.Chunks) {
|
||||
bz := []byte{}
|
||||
for _, chunk := range app.restoreChunks {
|
||||
bz = append(bz, chunk...)
|
||||
}
|
||||
err := app.state.Import(app.restoreSnapshot.Height, bz)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
app.restoreSnapshot = nil
|
||||
app.restoreChunks = nil
|
||||
}
|
||||
return abci.ResponseApplySnapshotChunk{Result: abci.ResponseApplySnapshotChunk_ACCEPT}
|
||||
}
|
||||
|
||||
// validatorUpdates generates a validator set update.
|
||||
func (app *Application) validatorUpdates(height uint64) (abci.ValidatorUpdates, error) {
|
||||
updates := app.cfg.ValidatorUpdates[fmt.Sprintf("%v", height)]
|
||||
if len(updates) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
valUpdates := abci.ValidatorUpdates{}
|
||||
for keyString, power := range updates {
|
||||
keyBytes, err := base64.StdEncoding.DecodeString(keyString)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid base64 pubkey value %q: %w", keyString, err)
|
||||
}
|
||||
valUpdates = append(valUpdates, abci.Ed25519ValidatorUpdate(keyBytes, int64(power)))
|
||||
}
|
||||
return valUpdates, nil
|
||||
}
|
||||
|
||||
// parseTx parses a tx in 'key=value' format into a key and value.
|
||||
func parseTx(tx []byte) (string, string, error) {
|
||||
parts := bytes.Split(tx, []byte("="))
|
||||
if len(parts) != 2 {
|
||||
return "", "", fmt.Errorf("invalid tx format: %q", string(tx))
|
||||
}
|
||||
if len(parts[0]) == 0 {
|
||||
return "", "", errors.New("key cannot be empty")
|
||||
}
|
||||
return string(parts[0]), string(parts[1]), nil
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/BurntSushi/toml"
|
||||
)
|
||||
|
||||
// Config is the application configuration.
|
||||
type Config struct {
|
||||
ChainID string `toml:"chain_id"`
|
||||
Listen string
|
||||
Protocol string
|
||||
Dir string
|
||||
PersistInterval uint64 `toml:"persist_interval"`
|
||||
SnapshotInterval uint64 `toml:"snapshot_interval"`
|
||||
RetainBlocks uint64 `toml:"retain_blocks"`
|
||||
ValidatorUpdates map[string]map[string]uint8 `toml:"validator_update"`
|
||||
PrivValServer string `toml:"privval_server"`
|
||||
PrivValKey string `toml:"privval_key"`
|
||||
PrivValState string `toml:"privval_state"`
|
||||
}
|
||||
|
||||
// LoadConfig loads the configuration from disk.
|
||||
func LoadConfig(file string) (*Config, error) {
|
||||
cfg := &Config{
|
||||
Listen: "unix:///var/run/app.sock",
|
||||
Protocol: "socket",
|
||||
PersistInterval: 1,
|
||||
}
|
||||
_, err := toml.DecodeFile(file, &cfg)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load config from %q: %w", file, err)
|
||||
}
|
||||
return cfg, cfg.Validate()
|
||||
}
|
||||
|
||||
// Validate validates the configuration. We don't do exhaustive config
|
||||
// validation here, instead relying on Testnet.Validate() to handle it.
|
||||
func (cfg Config) Validate() error {
|
||||
switch {
|
||||
case cfg.ChainID == "":
|
||||
return errors.New("chain_id parameter is required")
|
||||
case cfg.Listen == "" && cfg.Protocol != "builtin":
|
||||
return errors.New("listen parameter is required")
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/spf13/viper"
|
||||
"github.com/tendermint/tendermint/abci/server"
|
||||
"github.com/tendermint/tendermint/config"
|
||||
tmflags "github.com/tendermint/tendermint/libs/cli/flags"
|
||||
"github.com/tendermint/tendermint/libs/log"
|
||||
tmnet "github.com/tendermint/tendermint/libs/net"
|
||||
"github.com/tendermint/tendermint/node"
|
||||
"github.com/tendermint/tendermint/p2p"
|
||||
"github.com/tendermint/tendermint/privval"
|
||||
"github.com/tendermint/tendermint/proxy"
|
||||
)
|
||||
|
||||
var logger = log.NewTMLogger(log.NewSyncWriter(os.Stdout))
|
||||
|
||||
// main is the binary entrypoint.
|
||||
func main() {
|
||||
if len(os.Args) != 2 {
|
||||
fmt.Printf("Usage: %v <configfile>", os.Args[0])
|
||||
return
|
||||
}
|
||||
configFile := ""
|
||||
if len(os.Args) == 2 {
|
||||
configFile = os.Args[1]
|
||||
}
|
||||
|
||||
if err := run(configFile); err != nil {
|
||||
logger.Error(err.Error())
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
// run runs the application - basically like main() with error handling.
|
||||
func run(configFile string) error {
|
||||
cfg, err := LoadConfig(configFile)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
switch cfg.Protocol {
|
||||
case "socket", "grpc":
|
||||
err = startApp(cfg)
|
||||
case "builtin":
|
||||
err = startNode(cfg)
|
||||
default:
|
||||
err = fmt.Errorf("invalid protocol %q", cfg.Protocol)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Start remote signer
|
||||
if cfg.PrivValServer != "" {
|
||||
if err = startSigner(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// Apparently there's no way to wait for the server, so we just sleep
|
||||
for {
|
||||
time.Sleep(1 * time.Hour)
|
||||
}
|
||||
}
|
||||
|
||||
// startApp starts the application server, listening for connections from Tendermint.
|
||||
func startApp(cfg *Config) error {
|
||||
app, err := NewApplication(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
server, err := server.NewServer(cfg.Listen, cfg.Protocol, app)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = server.Start()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
logger.Info(fmt.Sprintf("Server listening on %v (%v protocol)", cfg.Listen, cfg.Protocol))
|
||||
return nil
|
||||
}
|
||||
|
||||
// startNode starts a Tendermint node running the application directly. It assumes the Tendermint
|
||||
// configuration is in $TMHOME/config/tendermint.toml.
|
||||
//
|
||||
// FIXME There is no way to simply load the configuration from a file, so we need to pull in Viper.
|
||||
func startNode(cfg *Config) error {
|
||||
app, err := NewApplication(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
home := os.Getenv("TMHOME")
|
||||
if home == "" {
|
||||
return errors.New("TMHOME not set")
|
||||
}
|
||||
viper.AddConfigPath(filepath.Join(home, "config"))
|
||||
viper.SetConfigName("config")
|
||||
err = viper.ReadInConfig()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmcfg := config.DefaultConfig()
|
||||
err = viper.Unmarshal(tmcfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmcfg.SetRoot(home)
|
||||
if err = tmcfg.ValidateBasic(); err != nil {
|
||||
return fmt.Errorf("error in config file: %v", err)
|
||||
}
|
||||
if tmcfg.LogFormat == config.LogFormatJSON {
|
||||
logger = log.NewTMJSONLogger(log.NewSyncWriter(os.Stdout))
|
||||
}
|
||||
logger, err = tmflags.ParseLogLevel(tmcfg.LogLevel, logger, config.DefaultLogLevel())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
logger = logger.With("module", "main")
|
||||
|
||||
nodeKey, err := p2p.LoadOrGenNodeKey(tmcfg.NodeKeyFile())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to load or gen node key %s: %w", tmcfg.NodeKeyFile(), err)
|
||||
}
|
||||
|
||||
n, err := node.NewNode(tmcfg,
|
||||
privval.LoadOrGenFilePV(tmcfg.PrivValidatorKeyFile(), tmcfg.PrivValidatorStateFile()),
|
||||
nodeKey,
|
||||
proxy.NewLocalClientCreator(app),
|
||||
node.DefaultGenesisDocProviderFunc(tmcfg),
|
||||
node.DefaultDBProvider,
|
||||
node.DefaultMetricsProvider(tmcfg.Instrumentation),
|
||||
logger,
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return n.Start()
|
||||
}
|
||||
|
||||
// startSigner starts a signer server connecting to the given endpoint.
|
||||
func startSigner(cfg *Config) error {
|
||||
filePV := privval.LoadFilePV(cfg.PrivValKey, cfg.PrivValState)
|
||||
|
||||
protocol, address := tmnet.ProtocolAndAddress(cfg.PrivValServer)
|
||||
var dialFn privval.SocketDialer
|
||||
switch protocol {
|
||||
case "tcp":
|
||||
dialFn = privval.DialTCPFn(address, 3*time.Second, filePV.Key.PrivKey)
|
||||
case "unix":
|
||||
dialFn = privval.DialUnixFn(address)
|
||||
default:
|
||||
return fmt.Errorf("invalid privval protocol %q", protocol)
|
||||
}
|
||||
|
||||
endpoint := privval.NewSignerDialerEndpoint(logger, dialFn,
|
||||
privval.SignerDialerEndpointRetryWaitInterval(1*time.Second),
|
||||
privval.SignerDialerEndpointConnRetries(100))
|
||||
err := privval.NewSignerServer(endpoint, cfg.ChainID, filePV).Start()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
logger.Info(fmt.Sprintf("Remote signer connecting to %v", cfg.PrivValServer))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
// nolint: gosec
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/ioutil"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
abci "github.com/tendermint/tendermint/abci/types"
|
||||
)
|
||||
|
||||
const (
|
||||
snapshotChunkSize = 1e6
|
||||
)
|
||||
|
||||
// SnapshotStore stores state sync snapshots. Snapshots are stored simply as
|
||||
// JSON files, and chunks are generated on-the-fly by splitting the JSON data
|
||||
// into fixed-size chunks.
|
||||
type SnapshotStore struct {
|
||||
sync.RWMutex
|
||||
dir string
|
||||
metadata []abci.Snapshot
|
||||
}
|
||||
|
||||
// NewSnapshotStore creates a new snapshot store.
|
||||
func NewSnapshotStore(dir string) (*SnapshotStore, error) {
|
||||
store := &SnapshotStore{dir: dir}
|
||||
if err := os.MkdirAll(dir, 0755); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := store.loadMetadata(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return store, nil
|
||||
}
|
||||
|
||||
// loadMetadata loads snapshot metadata. Does not take out locks, since it's
|
||||
// called internally on construction.
|
||||
func (s *SnapshotStore) loadMetadata() error {
|
||||
file := filepath.Join(s.dir, "metadata.json")
|
||||
metadata := []abci.Snapshot{}
|
||||
|
||||
bz, err := ioutil.ReadFile(file)
|
||||
switch {
|
||||
case errors.Is(err, os.ErrNotExist):
|
||||
case err != nil:
|
||||
return fmt.Errorf("failed to load snapshot metadata from %q: %w", file, err)
|
||||
}
|
||||
if len(bz) != 0 {
|
||||
err = json.Unmarshal(bz, &metadata)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid snapshot data in %q: %w", file, err)
|
||||
}
|
||||
}
|
||||
s.metadata = metadata
|
||||
return nil
|
||||
}
|
||||
|
||||
// saveMetadata saves snapshot metadata. Does not take out locks, since it's
|
||||
// called internally from e.g. Create().
|
||||
func (s *SnapshotStore) saveMetadata() error {
|
||||
bz, err := json.Marshal(s.metadata)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// save the file to a new file and move it to make saving atomic.
|
||||
newFile := filepath.Join(s.dir, "metadata.json.new")
|
||||
file := filepath.Join(s.dir, "metadata.json")
|
||||
err = ioutil.WriteFile(newFile, bz, 0644) // nolint: gosec
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(newFile, file)
|
||||
}
|
||||
|
||||
// Create creates a snapshot of the given application state's key/value pairs.
|
||||
func (s *SnapshotStore) Create(state *State) (abci.Snapshot, error) {
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
bz, err := state.Export()
|
||||
if err != nil {
|
||||
return abci.Snapshot{}, err
|
||||
}
|
||||
hash := sha256.Sum256(bz)
|
||||
snapshot := abci.Snapshot{
|
||||
Height: state.Height,
|
||||
Format: 1,
|
||||
Hash: hash[:],
|
||||
Chunks: byteChunks(bz),
|
||||
}
|
||||
err = ioutil.WriteFile(filepath.Join(s.dir, fmt.Sprintf("%v.json", state.Height)), bz, 0644)
|
||||
if err != nil {
|
||||
return abci.Snapshot{}, err
|
||||
}
|
||||
s.metadata = append(s.metadata, snapshot)
|
||||
err = s.saveMetadata()
|
||||
if err != nil {
|
||||
return abci.Snapshot{}, err
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
|
||||
// List lists available snapshots.
|
||||
func (s *SnapshotStore) List() ([]*abci.Snapshot, error) {
|
||||
s.RLock()
|
||||
defer s.RUnlock()
|
||||
snapshots := []*abci.Snapshot{}
|
||||
for _, snapshot := range s.metadata {
|
||||
s := snapshot // copy to avoid pointer to range variable
|
||||
snapshots = append(snapshots, &s)
|
||||
}
|
||||
return snapshots, nil
|
||||
}
|
||||
|
||||
// LoadChunk loads a snapshot chunk.
|
||||
func (s *SnapshotStore) LoadChunk(height uint64, format uint32, chunk uint32) ([]byte, error) {
|
||||
s.RLock()
|
||||
defer s.RUnlock()
|
||||
for _, snapshot := range s.metadata {
|
||||
if snapshot.Height == height && snapshot.Format == format {
|
||||
bz, err := ioutil.ReadFile(filepath.Join(s.dir, fmt.Sprintf("%v.json", height)))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return byteChunk(bz, chunk), nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// byteChunk returns the chunk at a given index from the full byte slice.
|
||||
func byteChunk(bz []byte, index uint32) []byte {
|
||||
start := int(index * snapshotChunkSize)
|
||||
end := int((index + 1) * snapshotChunkSize)
|
||||
switch {
|
||||
case start >= len(bz):
|
||||
return nil
|
||||
case end >= len(bz):
|
||||
return bz[start:]
|
||||
default:
|
||||
return bz[start:end]
|
||||
}
|
||||
}
|
||||
|
||||
// byteChunks calculates the number of chunks in the byte slice.
|
||||
func byteChunks(bz []byte) uint32 {
|
||||
return uint32(math.Ceil(float64(len(bz)) / snapshotChunkSize))
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
//nolint: gosec
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/ioutil"
|
||||
"os"
|
||||
"sort"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// State is the application state.
|
||||
type State struct {
|
||||
sync.RWMutex
|
||||
Height uint64
|
||||
Values map[string]string
|
||||
Hash []byte
|
||||
|
||||
// private fields aren't marshalled to disk.
|
||||
file string
|
||||
persistInterval uint64
|
||||
initialHeight uint64
|
||||
}
|
||||
|
||||
// NewState creates a new state.
|
||||
func NewState(file string, persistInterval uint64) (*State, error) {
|
||||
state := &State{
|
||||
Values: make(map[string]string),
|
||||
file: file,
|
||||
persistInterval: persistInterval,
|
||||
}
|
||||
state.Hash = hashItems(state.Values)
|
||||
err := state.load()
|
||||
switch {
|
||||
case errors.Is(err, os.ErrNotExist):
|
||||
case err != nil:
|
||||
return nil, err
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
|
||||
// load loads state from disk. It does not take out a lock, since it is called
|
||||
// during construction.
|
||||
func (s *State) load() error {
|
||||
bz, err := ioutil.ReadFile(s.file)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read state from %q: %w", s.file, err)
|
||||
}
|
||||
err = json.Unmarshal(bz, s)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid state data in %q: %w", s.file, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// save saves the state to disk. It does not take out a lock since it is called
|
||||
// internally by Commit which does lock.
|
||||
func (s *State) save() error {
|
||||
bz, err := json.Marshal(s)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal state: %w", err)
|
||||
}
|
||||
// We write the state to a separate file and move it to the destination, to
|
||||
// make it atomic.
|
||||
newFile := fmt.Sprintf("%v.new", s.file)
|
||||
err = ioutil.WriteFile(newFile, bz, 0644)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to write state to %q: %w", s.file, err)
|
||||
}
|
||||
return os.Rename(newFile, s.file)
|
||||
}
|
||||
|
||||
// Export exports key/value pairs as JSON, used for state sync snapshots.
|
||||
func (s *State) Export() ([]byte, error) {
|
||||
s.RLock()
|
||||
defer s.RUnlock()
|
||||
return json.Marshal(s.Values)
|
||||
}
|
||||
|
||||
// Import imports key/value pairs from JSON bytes, used for InitChain.AppStateBytes and
|
||||
// state sync snapshots. It also saves the state once imported.
|
||||
func (s *State) Import(height uint64, jsonBytes []byte) error {
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
values := map[string]string{}
|
||||
err := json.Unmarshal(jsonBytes, &values)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to decode imported JSON data: %w", err)
|
||||
}
|
||||
s.Height = height
|
||||
s.Values = values
|
||||
s.Hash = hashItems(values)
|
||||
return s.save()
|
||||
}
|
||||
|
||||
// Get fetches a value. A missing value is returned as an empty string.
|
||||
func (s *State) Get(key string) string {
|
||||
s.RLock()
|
||||
defer s.RUnlock()
|
||||
return s.Values[key]
|
||||
}
|
||||
|
||||
// Set sets a value. Setting an empty value is equivalent to deleting it.
|
||||
func (s *State) Set(key, value string) {
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
if value == "" {
|
||||
delete(s.Values, key)
|
||||
} else {
|
||||
s.Values[key] = value
|
||||
}
|
||||
}
|
||||
|
||||
// Commit commits the current state.
|
||||
func (s *State) Commit() (uint64, []byte, error) {
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
s.Hash = hashItems(s.Values)
|
||||
switch {
|
||||
case s.Height > 0:
|
||||
s.Height++
|
||||
case s.initialHeight > 0:
|
||||
s.Height = s.initialHeight
|
||||
default:
|
||||
s.Height = 1
|
||||
}
|
||||
if s.persistInterval > 0 && s.Height%s.persistInterval == 0 {
|
||||
err := s.save()
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
}
|
||||
return s.Height, s.Hash, nil
|
||||
}
|
||||
|
||||
// hashItems hashes a set of key/value items.
|
||||
func hashItems(items map[string]string) []byte {
|
||||
keys := make([]string, 0, len(items))
|
||||
for key := range items {
|
||||
keys = append(keys, key)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
|
||||
hasher := sha256.New()
|
||||
for _, key := range keys {
|
||||
_, _ = hasher.Write([]byte(key))
|
||||
_, _ = hasher.Write([]byte{0})
|
||||
_, _ = hasher.Write([]byte(items[key]))
|
||||
_, _ = hasher.Write([]byte{0})
|
||||
}
|
||||
return hasher.Sum(nil)
|
||||
}
|
||||
Reference in New Issue
Block a user