← Return to the article

Companion code

The complete 4+2 demo.

Dependency-free teaching code for the exact codec used in the article. It is intentionally scalar and explicit. It is not a production storage library.

field.goGF(2⁸) multiplication, powers, and inverses.
package erasure

// This demo uses GF(2^8) with the primitive polynomial x^8+x^4+x^3+x^2+1
// (0x11d). Addition and subtraction are XOR.

func gfMul(a, b byte) byte {
	var product uint16
	left := uint16(a)
	right := uint16(b)

	for right != 0 {
		if right&1 != 0 {
			product ^= left
		}
		right >>= 1
		left <<= 1
		if left&0x100 != 0 {
			left ^= 0x11d
		}
	}

	return byte(product)
}

func gfPow(base byte, exponent int) byte {
	result := byte(1)
	for exponent > 0 {
		if exponent&1 != 0 {
			result = gfMul(result, base)
		}
		base = gfMul(base, base)
		exponent >>= 1
	}
	return result
}

func gfInv(value byte) byte {
	if value == 0 {
		panic("zero has no multiplicative inverse")
	}
	// In a field with 256 elements, a^255 = 1 for every nonzero a.
	return gfPow(value, 254)
}
matrix.goCauchy generator construction and Gauss-Jordan inversion.
package erasure

import "fmt"

type matrix [][]byte

func makeMatrix(rows, columns int) matrix {
	result := make(matrix, rows)
	for row := range result {
		result[row] = make([]byte, columns)
	}
	return result
}

func identityMatrix(size int) matrix {
	result := makeMatrix(size, size)
	for i := 0; i < size; i++ {
		result[i][i] = 1
	}
	return result
}

func (m matrix) clone() matrix {
	result := makeMatrix(len(m), len(m[0]))
	for row := range m {
		copy(result[row], m[row])
	}
	return result
}

func multiplyMatrices(left, right matrix) (matrix, error) {
	if len(left) == 0 || len(right) == 0 || len(left[0]) != len(right) {
		return nil, fmt.Errorf("incompatible matrix dimensions")
	}

	result := makeMatrix(len(left), len(right[0]))
	for row := range left {
		for column := range right[0] {
			var value byte
			for i := range right {
				value ^= gfMul(left[row][i], right[i][column])
			}
			result[row][column] = value
		}
	}
	return result, nil
}

func invertMatrix(input matrix) (matrix, error) {
	size := len(input)
	if size == 0 {
		return nil, fmt.Errorf("cannot invert an empty matrix")
	}
	for _, row := range input {
		if len(row) != size {
			return nil, fmt.Errorf("matrix must be square")
		}
	}

	left := input.clone()
	right := identityMatrix(size)

	for column := 0; column < size; column++ {
		if left[column][column] == 0 {
			swap := column + 1
			for swap < size && left[swap][column] == 0 {
				swap++
			}
			if swap == size {
				return nil, fmt.Errorf("matrix is singular")
			}
			left[column], left[swap] = left[swap], left[column]
			right[column], right[swap] = right[swap], right[column]
		}

		pivot := left[column][column]
		if pivot != 1 {
			scale := gfInv(pivot)
			for i := 0; i < size; i++ {
				left[column][i] = gfMul(left[column][i], scale)
				right[column][i] = gfMul(right[column][i], scale)
			}
		}

		for row := 0; row < size; row++ {
			if row == column || left[row][column] == 0 {
				continue
			}
			factor := left[row][column]
			for i := 0; i < size; i++ {
				left[row][i] ^= gfMul(factor, left[column][i])
				right[row][i] ^= gfMul(factor, right[column][i])
			}
		}
	}

	return right, nil
}

// cauchyGenerator returns a systematic totalShards × dataShards matrix.
// The identity rows preserve the original data. The remaining rows use the
// same inverse(i XOR j) construction exposed by ISA-L's teaching-level base
// implementation.
func cauchyGenerator(dataShards, totalShards int) matrix {
	result := makeMatrix(totalShards, dataShards)
	for i := 0; i < dataShards; i++ {
		result[i][i] = 1
	}
	for row := dataShards; row < totalShards; row++ {
		for column := 0; column < dataShards; column++ {
			result[row][column] = gfInv(byte(row ^ column))
		}
	}
	return result
}
codec.goSystematic 4+2 encoding, checksums, reconstruction, and joining.
package erasure

import (
	"bytes"
	"encoding/json"
	"errors"
	"fmt"
	"hash/crc32"
)

const (
	DataShards   = 4
	ParityShards = 2
	TotalShards  = DataShards + ParityShards
	CodecVersion = 1
)

var (
	ErrTooFewShards = errors.New("fewer than four valid shards remain")
	ErrCorruptShard = errors.New("a surviving shard failed its checksum")
)

type Metadata struct {
	OriginalSize int                 `json:"original_size"`
	ShardSize    int                 `json:"shard_size"`
	CodecVersion int                 `json:"codec_version"`
	Checksums    [TotalShards]uint32 `json:"checksums"`
}

type Codec struct {
	generator matrix
}

func New4Plus2() *Codec {
	return &Codec{generator: cauchyGenerator(DataShards, TotalShards)}
}

func (c *Codec) Encode(payload []byte) ([][]byte, Metadata, error) {
	shardSize := (len(payload) + DataShards - 1) / DataShards
	if shardSize == 0 {
		shardSize = 1
	}

	shards := make([][]byte, TotalShards)
	for i := range shards {
		shards[i] = make([]byte, shardSize)
	}

	for i := 0; i < DataShards; i++ {
		start := i * shardSize
		if start >= len(payload) {
			break
		}
		end := min(start+shardSize, len(payload))
		copy(shards[i], payload[start:end])
	}

	for row := DataShards; row < TotalShards; row++ {
		encodeRow(shards[row], c.generator[row], shards[:DataShards])
	}

	meta := Metadata{
		OriginalSize: len(payload),
		ShardSize:    shardSize,
		CodecVersion: CodecVersion,
	}
	refreshChecksums(shards, &meta)
	return shards, meta, nil
}

func (c *Codec) Reconstruct(shards [][]byte, meta *Metadata) error {
	if err := validateMetadata(meta); err != nil {
		return err
	}
	if len(shards) != TotalShards {
		return fmt.Errorf("need exactly %d shard slots", TotalShards)
	}

	present := make([]int, 0, DataShards)
	for index, shard := range shards {
		if shard == nil {
			continue
		}
		if len(shard) != meta.ShardSize {
			return fmt.Errorf("shard %d has length %d, want %d", index, len(shard), meta.ShardSize)
		}
		if crc32.ChecksumIEEE(shard) != meta.Checksums[index] {
			return fmt.Errorf("%w: shard %d", ErrCorruptShard, index)
		}
		if len(present) < DataShards {
			present = append(present, index)
		}
	}
	if len(present) < DataShards {
		return ErrTooFewShards
	}

	survivorMatrix := makeMatrix(DataShards, DataShards)
	survivorShards := make([][]byte, DataShards)
	for row, shardIndex := range present {
		copy(survivorMatrix[row], c.generator[shardIndex])
		survivorShards[row] = shards[shardIndex]
	}

	decodeMatrix, err := invertMatrix(survivorMatrix)
	if err != nil {
		return fmt.Errorf("survivor matrix: %w", err)
	}

	data := make([][]byte, DataShards)
	for row := 0; row < DataShards; row++ {
		data[row] = make([]byte, meta.ShardSize)
		encodeRow(data[row], decodeMatrix[row], survivorShards)
	}

	for row := 0; row < DataShards; row++ {
		if shards[row] == nil {
			shards[row] = data[row]
		}
	}
	for row := DataShards; row < TotalShards; row++ {
		if shards[row] == nil {
			shards[row] = make([]byte, meta.ShardSize)
			encodeRow(shards[row], c.generator[row], data)
		}
	}

	refreshChecksums(shards, meta)
	return nil
}

func Join(shards [][]byte, meta Metadata) ([]byte, error) {
	if err := validateMetadata(&meta); err != nil {
		return nil, err
	}
	if len(shards) != TotalShards {
		return nil, fmt.Errorf("need exactly %d shard slots", TotalShards)
	}

	var padded bytes.Buffer
	for index := 0; index < DataShards; index++ {
		if shards[index] == nil || len(shards[index]) != meta.ShardSize {
			return nil, fmt.Errorf("data shard %d is absent or has the wrong size", index)
		}
		padded.Write(shards[index])
	}
	if meta.OriginalSize > padded.Len() {
		return nil, fmt.Errorf("original size exceeds padded data")
	}
	return bytes.Clone(padded.Bytes()[:meta.OriginalSize]), nil
}

func (meta Metadata) JSON() ([]byte, error) {
	return json.MarshalIndent(meta, "", "  ")
}

func MetadataFromJSON(data []byte) (Metadata, error) {
	var meta Metadata
	if err := json.Unmarshal(data, &meta); err != nil {
		return Metadata{}, err
	}
	return meta, validateMetadata(&meta)
}

func encodeRow(output []byte, coefficients []byte, inputs [][]byte) {
	for inputIndex, input := range inputs {
		coefficient := coefficients[inputIndex]
		for byteIndex, value := range input {
			output[byteIndex] ^= gfMul(coefficient, value)
		}
	}
}

func refreshChecksums(shards [][]byte, meta *Metadata) {
	for index, shard := range shards {
		meta.Checksums[index] = crc32.ChecksumIEEE(shard)
	}
}

func validateMetadata(meta *Metadata) error {
	if meta == nil {
		return fmt.Errorf("metadata is required")
	}
	if meta.CodecVersion != CodecVersion {
		return fmt.Errorf("unsupported codec version %d", meta.CodecVersion)
	}
	if meta.ShardSize < 1 || meta.OriginalSize < 0 || meta.OriginalSize > DataShards*meta.ShardSize {
		return fmt.Errorf("invalid size metadata")
	}
	return nil
}
codec_test.goField identities, every survivor matrix, all 15 two-shard failures, and error cases.
package erasure

import (
	"bytes"
	"errors"
	"fmt"
	"testing"
)

func TestFieldIdentities(t *testing.T) {
	for value := 0; value < 256; value++ {
		a := byte(value)
		if a^0 != a || a^a != 0 || gfMul(a, 1) != a {
			t.Fatalf("field identity failed for %d", value)
		}
		if a != 0 && gfMul(a, gfInv(a)) != 1 {
			t.Fatalf("inverse failed for %d", value)
		}
	}
}

func TestGeneratorIsSystematicAndMDS(t *testing.T) {
	generator := cauchyGenerator(DataShards, TotalShards)
	for row := 0; row < DataShards; row++ {
		for column := 0; column < DataShards; column++ {
			want := byte(0)
			if row == column {
				want = 1
			}
			if generator[row][column] != want {
				t.Fatalf("generator[%d][%d] = %d, want %d", row, column, generator[row][column], want)
			}
		}
	}

	for firstMissing := 0; firstMissing < TotalShards; firstMissing++ {
		for secondMissing := firstMissing + 1; secondMissing < TotalShards; secondMissing++ {
			survivors := makeMatrix(DataShards, DataShards)
			row := 0
			for index := 0; index < TotalShards; index++ {
				if index == firstMissing || index == secondMissing {
					continue
				}
				copy(survivors[row], generator[index])
				row++
			}
			if _, err := invertMatrix(survivors); err != nil {
				t.Fatalf("missing (%d,%d) produced a singular survivor matrix: %v", firstMissing, secondMissing, err)
			}
		}
	}
}

func TestEveryTwoShardLossRecoversEveryByte(t *testing.T) {
	payload := []byte("Erasure coding does not save the missing bytes. It saves enough independent equations to calculate them again.")
	codec := New4Plus2()
	originalShards, originalMeta, err := codec.Encode(payload)
	if err != nil {
		t.Fatal(err)
	}

	patterns := 0
	for firstMissing := 0; firstMissing < TotalShards; firstMissing++ {
		for secondMissing := firstMissing + 1; secondMissing < TotalShards; secondMissing++ {
			patterns++
			shards := cloneShards(originalShards)
			meta := originalMeta
			shards[firstMissing] = nil
			shards[secondMissing] = nil

			if err := codec.Reconstruct(shards, &meta); err != nil {
				t.Fatalf("recover missing (%d,%d): %v", firstMissing, secondMissing, err)
			}
			got, err := Join(shards, meta)
			if err != nil {
				t.Fatal(err)
			}
			if !bytes.Equal(got, payload) {
				t.Fatalf("missing (%d,%d): recovered bytes differ", firstMissing, secondMissing)
			}
			for index := range shards {
				if !bytes.Equal(shards[index], originalShards[index]) {
					t.Fatalf("missing (%d,%d): shard %d was not regenerated identically", firstMissing, secondMissing, index)
				}
			}
		}
	}
	if patterns != 15 {
		t.Fatalf("tested %d patterns, want 15", patterns)
	}
}

func TestLengthsAroundShardBoundaries(t *testing.T) {
	codec := New4Plus2()
	for size := 0; size <= 17; size++ {
		t.Run(fmt.Sprintf("size-%d", size), func(t *testing.T) {
			payload := make([]byte, size)
			for index := range payload {
				payload[index] = byte(index * 17)
			}
			shards, meta, err := codec.Encode(payload)
			if err != nil {
				t.Fatal(err)
			}
			shards[0], shards[5] = nil, nil
			if err := codec.Reconstruct(shards, &meta); err != nil {
				t.Fatal(err)
			}
			got, err := Join(shards, meta)
			if err != nil {
				t.Fatal(err)
			}
			if !bytes.Equal(got, payload) {
				t.Fatalf("round trip differs: got %v want %v", got, payload)
			}
		})
	}
}

func TestThreeShardLossFails(t *testing.T) {
	codec := New4Plus2()
	shards, meta, _ := codec.Encode([]byte("only three survivors remain"))
	shards[0], shards[2], shards[5] = nil, nil, nil

	if err := codec.Reconstruct(shards, &meta); !errors.Is(err, ErrTooFewShards) {
		t.Fatalf("got %v, want ErrTooFewShards", err)
	}
}

func TestCorruptSurvivorFailsChecksum(t *testing.T) {
	codec := New4Plus2()
	shards, meta, _ := codec.Encode([]byte("checksums turn corruption into a known failure"))
	shards[3][0] ^= 0xff

	if err := codec.Reconstruct(shards, &meta); !errors.Is(err, ErrCorruptShard) {
		t.Fatalf("got %v, want ErrCorruptShard", err)
	}
}

func cloneShards(input [][]byte) [][]byte {
	result := make([][]byte, len(input))
	for index := range input {
		result[index] = bytes.Clone(input[index])
	}
	return result
}
cmd/erasure-demo/main.goThe prove, encode, lose, and recover commands.
package main

import (
	"fmt"
	"log"
	"os"
	"strconv"

	erasure "harshitsharma.co/examples/erasure-coding"
)

func main() {
	if len(os.Args) < 2 {
		usage()
	}

	switch os.Args[1] {
	case "prove":
		prove()
	case "encode":
		if len(os.Args) != 4 {
			usage()
		}
		encode(os.Args[2], os.Args[3])
	case "lose":
		if len(os.Args) < 4 {
			usage()
		}
		lose(os.Args[2], os.Args[3:])
	case "recover":
		if len(os.Args) != 4 {
			usage()
		}
		recover(os.Args[2], os.Args[3])
	default:
		usage()
	}
}

func prove() {
	payload := []byte("Every byte returns after every possible two-shard loss.")
	codec := erasure.New4Plus2()
	encoded, meta, err := codec.Encode(payload)
	must(err)

	patterns := 0
	for first := 0; first < erasure.TotalShards; first++ {
		for second := first + 1; second < erasure.TotalShards; second++ {
			shards := clone(encoded)
			shards[first], shards[second] = nil, nil
			must(codec.Reconstruct(shards, &meta))
			recovered, err := erasure.Join(shards, meta)
			must(err)
			if string(recovered) != string(payload) {
				log.Fatalf("pattern (%d,%d) reconstructed different bytes", first, second)
			}
			patterns++
		}
	}
	fmt.Printf("verified all %d two-shard loss patterns; every byte recovered\n", patterns)
}

func encode(inputPath, directory string) {
	payload, err := os.ReadFile(inputPath)
	must(err)
	shards, meta, err := erasure.New4Plus2().Encode(payload)
	must(err)
	must(os.MkdirAll(directory, 0o755))
	for index, shard := range shards {
		must(os.WriteFile(shardPath(directory, index), shard, 0o644))
	}
	data, err := meta.JSON()
	must(err)
	must(os.WriteFile(directory+"/metadata.json", data, 0o644))
	fmt.Printf("encoded %d bytes into six %d-byte shards\n", len(payload), meta.ShardSize)
}

func lose(directory string, values []string) {
	for _, value := range values {
		index, err := strconv.Atoi(value)
		must(err)
		if index < 0 || index >= erasure.TotalShards {
			log.Fatalf("shard index %d is outside 0..5", index)
		}
		path := shardPath(directory, index)
		must(os.Rename(path, path+".missing"))
		fmt.Printf("moved shard %d aside as %s.missing\n", index, path)
	}
}

func recover(directory, outputPath string) {
	metaData, err := os.ReadFile(directory + "/metadata.json")
	must(err)
	meta, err := erasure.MetadataFromJSON(metaData)
	must(err)

	shards := make([][]byte, erasure.TotalShards)
	for index := range shards {
		data, readErr := os.ReadFile(shardPath(directory, index))
		if readErr == nil {
			shards[index] = data
			continue
		}
		if !os.IsNotExist(readErr) {
			must(readErr)
		}
	}

	must(erasure.New4Plus2().Reconstruct(shards, &meta))
	payload, err := erasure.Join(shards, meta)
	must(err)
	must(os.WriteFile(outputPath, payload, 0o644))
	fmt.Printf("recovered every byte to %s\n", outputPath)
}

func shardPath(directory string, index int) string {
	return fmt.Sprintf("%s/shard-%d.bin", directory, index)
}

func clone(input [][]byte) [][]byte {
	result := make([][]byte, len(input))
	for index := range input {
		result[index] = append([]byte(nil), input[index]...)
	}
	return result
}

func must(err error) {
	if err != nil {
		log.Fatal(err)
	}
}

func usage() {
	fmt.Fprintln(os.Stderr, `usage:
  erasure-demo prove
  erasure-demo encode INPUT DIRECTORY
  erasure-demo lose DIRECTORY SHARD_INDEX [SHARD_INDEX...]
  erasure-demo recover DIRECTORY OUTPUT`)
	os.Exit(2)
}