121 lines
3.0 KiB
Go
121 lines
3.0 KiB
Go
package state
|
|
|
|
import (
|
|
"bytes"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"io"
|
|
"os"
|
|
)
|
|
|
|
const maxSnapshotBytes = 16 << 20
|
|
|
|
type snapshot struct {
|
|
Version int `json:"version"`
|
|
HostID string `json:"hostId"`
|
|
Operations map[string]Operation `json:"operations"`
|
|
Checksum string `json:"checksum"`
|
|
}
|
|
|
|
func encodeSnapshot(s snapshot) ([]byte, error) {
|
|
s.Checksum = ""
|
|
raw, err := json.Marshal(s)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
sum := sha256.Sum256(raw)
|
|
s.Checksum = "sha256:" + hex.EncodeToString(sum[:])
|
|
raw, err = json.Marshal(s)
|
|
if len(raw) > maxSnapshotBytes {
|
|
return nil, errors.New("state capacity exceeded")
|
|
}
|
|
return raw, err
|
|
}
|
|
|
|
func readSnapshot(r *os.Root, hostID string, initialized bool) (snapshot, error) {
|
|
empty := snapshot{Version: 1, HostID: hostID, Operations: make(map[string]Operation)}
|
|
if err := regularOrMissing(r, "state.json"); err != nil {
|
|
return snapshot{}, err
|
|
}
|
|
f, err := r.Open("state.json")
|
|
if os.IsNotExist(err) {
|
|
if initialized {
|
|
return snapshot{}, errors.New("initialized store has lost its snapshot")
|
|
}
|
|
return empty, nil
|
|
}
|
|
if err != nil {
|
|
return snapshot{}, err
|
|
}
|
|
defer f.Close()
|
|
raw, err := io.ReadAll(io.LimitReader(f, maxSnapshotBytes+1))
|
|
if err != nil {
|
|
return snapshot{}, err
|
|
}
|
|
if len(raw) > maxSnapshotBytes {
|
|
return snapshot{}, errors.New("state exceeds size limit")
|
|
}
|
|
var s snapshot
|
|
if err := json.Unmarshal(raw, &s); err != nil {
|
|
return snapshot{}, errors.New("corrupt state")
|
|
}
|
|
if s.Version != 1 || s.HostID != hostID || s.Operations == nil {
|
|
return snapshot{}, errors.New("incompatible state or host mismatch")
|
|
}
|
|
canonical, err := encodeSnapshot(s)
|
|
if err != nil || !bytes.Equal(raw, canonical) {
|
|
return snapshot{}, errors.New("state integrity check failed")
|
|
}
|
|
unresolved := 0
|
|
for id, op := range s.Operations {
|
|
if id != op.ID || !idPattern.MatchString(id) || !hashPattern.MatchString(op.PlanHash) || !op.Status.valid() || op.Revision == 0 {
|
|
return snapshot{}, errors.New("invalid operation record")
|
|
}
|
|
if !op.Status.terminal() {
|
|
unresolved++
|
|
}
|
|
}
|
|
if unresolved > 1 {
|
|
return snapshot{}, errors.New("multiple unresolved operations")
|
|
}
|
|
return s, nil
|
|
}
|
|
|
|
func writeSnapshot(r *os.Root, s snapshot) error {
|
|
raw, err := encodeSnapshot(s)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if err := regularOrMissing(r, "state.json"); err != nil {
|
|
return err
|
|
}
|
|
var entropy [16]byte
|
|
if _, err := rand.Read(entropy[:]); err != nil {
|
|
return err
|
|
}
|
|
name := ".state-" + hex.EncodeToString(entropy[:]) + ".tmp"
|
|
f, err := r.OpenFile(name, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0600)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer r.Remove(name) // Only the unique file just created; never state.json or host.lock.
|
|
if _, err := f.Write(raw); err != nil {
|
|
f.Close()
|
|
return err
|
|
}
|
|
if err := f.Sync(); err != nil {
|
|
f.Close()
|
|
return err
|
|
}
|
|
if err := f.Close(); err != nil {
|
|
return err
|
|
}
|
|
if err := r.Rename(name, "state.json"); err != nil {
|
|
return err
|
|
}
|
|
return syncDirectory(r)
|
|
}
|