Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 46 additions & 0 deletions market/matching/delta_validation_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
package matching

import (
"testing"

"github.com/shopspring/decimal"
"github.com/tent-of-trials/market/orderbook"
"github.com/tent-of-trials/market/types"
)

func TestMatchingBookPreservesStateAfterInvalidDelta(t *testing.T) {
book := orderbook.NewOrderBook("BTC-USDC", orderbook.Config{MaxDepth: 100})
engine := NewMatchingEngine(EngineConfig{EnableShorting: true}, map[types.Symbol]*orderbook.OrderBook{
"BTC-USDC": book,
})
_, err := engine.PlaceOrder(&types.Order{
Symbol: "BTC-USDC",
Side: types.Buy,
Type: types.Limit,
Price: decimal.RequireFromString("100"),
Quantity: decimal.RequireFromString("1"),
RemainingQty: decimal.RequireFromString("1"),
})
if err != nil {
t.Fatalf("place order: %v", err)
}
before := book.GetSnapshot()

err = book.ApplyDelta(orderbook.Delta{
Symbol: "BTC-USDC",
Sequence: book.Sequence() + 1,
Bids: []types.Level{{
Price: decimal.RequireFromString("99"),
Quantity: decimal.RequireFromString("-1"),
Count: 1,
}},
})
if err == nil {
t.Fatal("expected invalid quantity error")
}

after := book.GetSnapshot()
if len(after.Bids) != len(before.Bids) || !after.Bids[0].Quantity.Equal(before.Bids[0].Quantity) {
t.Fatalf("book mutated after invalid delta: before %+v after %+v", before, after)
}
}
159 changes: 159 additions & 0 deletions market/orderbook/delta.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
package orderbook

import (
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"

"github.com/shopspring/decimal"
"github.com/tent-of-trials/market/types"
)

var (
ErrInvalidDeltaSymbol = errors.New("invalid delta symbol")
ErrInvalidDeltaSequence = errors.New("stale or out-of-order delta sequence")
ErrInvalidDeltaPrice = errors.New("invalid delta price")
ErrInvalidDeltaQuantity = errors.New("invalid delta quantity")
ErrInvalidDeltaChecksum = errors.New("invalid delta checksum")
)

type Delta struct {
Symbol types.Symbol `json:"symbol"`
Sequence uint64 `json:"sequence"`
Bids []types.Level `json:"bids"`
Asks []types.Level `json:"asks"`
Checksum string `json:"checksum,omitempty"`
}

func (ob *OrderBook) Sequence() uint64 {
ob.mu.RLock()
defer ob.mu.RUnlock()
return ob.sequence
}

func (ob *OrderBook) ReplaceSnapshot(snapshot *types.DepthUpdate, sequence uint64) error {
if snapshot == nil || snapshot.Symbol == "" || snapshot.Symbol != ob.symbol {
return ErrInvalidDeltaSymbol
}
if err := validateLevels(snapshot.Bids, "bid"); err != nil {
return err
}
if err := validateLevels(snapshot.Asks, "ask"); err != nil {
return err
}

ob.mu.Lock()
defer ob.mu.Unlock()
if ob.closed {
return ErrBookClosed
}
ob.bids = cloneLevelPointers(snapshot.Bids)
ob.asks = cloneLevelPointers(snapshot.Asks)
sortLevels(ob.bids, true)
sortLevels(ob.asks, false)
ob.sequence = sequence
ob.updatedAt = time.Now()
return nil
}

func (ob *OrderBook) ApplyDelta(delta Delta) error {
if delta.Symbol == "" || delta.Symbol != ob.symbol {
return ErrInvalidDeltaSymbol
}
if err := validateLevels(delta.Bids, "bid"); err != nil {
return err
}
if err := validateLevels(delta.Asks, "ask"); err != nil {
return err
}

ob.mu.Lock()
defer ob.mu.Unlock()
if ob.closed {
return ErrBookClosed
}
if delta.Sequence <= ob.sequence {
return ErrInvalidDeltaSequence
}

nextBids := cloneLevelPointersFromPointers(ob.bids)
nextAsks := cloneLevelPointersFromPointers(ob.asks)
nextBids = applyLevels(nextBids, delta.Bids, true)
nextAsks = applyLevels(nextAsks, delta.Asks, false)

if delta.Checksum != "" {
got := ComputeChecksum(nextBids, nextAsks)
if !strings.EqualFold(delta.Checksum, got) {
return fmt.Errorf("%w: expected %s got %s", ErrInvalidDeltaChecksum, delta.Checksum, got)
}
}

ob.bids = nextBids
ob.asks = nextAsks
ob.sequence = delta.Sequence
ob.updatedAt = time.Now()
return nil
}

func ComputeChecksum(bids []*types.Level, asks []*types.Level) string {
h := sha256.New()
writeLevels := func(side string, levels []*types.Level) {
for _, level := range levels {
if level == nil {
continue
}
fmt.Fprintf(h, "%s:%s:%s:%d;", side, level.Price.String(), level.Quantity.String(), level.Count)
}
}
writeLevels("bid", bids)
writeLevels("ask", asks)
return hex.EncodeToString(h.Sum(nil))
}

func validateLevels(levels []types.Level, side string) error {
for i, level := range levels {
if level.Price.LessThanOrEqual(decimal.Zero) {
return fmt.Errorf("%w: %s[%d]", ErrInvalidDeltaPrice, side, i)
}
if level.Quantity.LessThan(decimal.Zero) {
return fmt.Errorf("%w: %s[%d]", ErrInvalidDeltaQuantity, side, i)
}
}
return nil
}

func applyLevels(current []*types.Level, updates []types.Level, desc bool) []*types.Level {
for _, update := range updates {
current = removeLevel(current, update.Price)
if update.Quantity.GreaterThan(decimal.Zero) {
level := update
current = append(current, &level)
}
}
sortLevels(current, desc)
return current
}

func cloneLevelPointers(levels []types.Level) []*types.Level {
result := make([]*types.Level, 0, len(levels))
for _, level := range levels {
copy := level
result = append(result, &copy)
}
return result
}

func cloneLevelPointersFromPointers(levels []*types.Level) []*types.Level {
result := make([]*types.Level, 0, len(levels))
for _, level := range levels {
if level == nil {
continue
}
copy := *level
result = append(result, &copy)
}
return result
}
182 changes: 182 additions & 0 deletions market/orderbook/delta_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,182 @@
package orderbook

import (
"errors"
"testing"

"github.com/shopspring/decimal"
"github.com/tent-of-trials/market/types"
)

func TestApplyDeltaRejectsMalformedLevelsWithoutMutatingState(t *testing.T) {
book := seededBook(t)
before := book.GetSnapshot()

err := book.ApplyDelta(Delta{
Symbol: "BTC-USDC",
Sequence: book.Sequence() + 1,
Bids: []types.Level{{
Price: decimal.RequireFromString("-1"),
Quantity: decimal.RequireFromString("0.5"),
Count: 1,
}},
})

if !errors.Is(err, ErrInvalidDeltaPrice) {
t.Fatalf("expected invalid price error, got %v", err)
}
assertSnapshotEqual(t, before, book.GetSnapshot())
}

func TestApplyDeltaRejectsStaleSequenceWithoutMutatingState(t *testing.T) {
book := seededBook(t)
before := book.GetSnapshot()

err := book.ApplyDelta(Delta{
Symbol: "BTC-USDC",
Sequence: book.Sequence(),
Asks: []types.Level{{
Price: decimal.RequireFromString("103"),
Quantity: decimal.RequireFromString("1"),
Count: 1,
}},
})

if !errors.Is(err, ErrInvalidDeltaSequence) {
t.Fatalf("expected stale sequence error, got %v", err)
}
assertSnapshotEqual(t, before, book.GetSnapshot())
}

func TestApplyDeltaRejectsChecksumMismatchWithoutMutatingState(t *testing.T) {
book := seededBook(t)
before := book.GetSnapshot()

err := book.ApplyDelta(Delta{
Symbol: "BTC-USDC",
Sequence: book.Sequence() + 1,
Asks: []types.Level{{
Price: decimal.RequireFromString("103"),
Quantity: decimal.RequireFromString("1"),
Count: 1,
}},
Checksum: "not-the-next-book-checksum",
})

if !errors.Is(err, ErrInvalidDeltaChecksum) {
t.Fatalf("expected checksum error, got %v", err)
}
assertSnapshotEqual(t, before, book.GetSnapshot())
}

func TestReplaceSnapshotThenApplyValidDelta(t *testing.T) {
book := NewOrderBook("ETH-USDC", Config{MaxDepth: 100})
err := book.ReplaceSnapshot(&types.DepthUpdate{
Symbol: "ETH-USDC",
Bids: []types.Level{{
Price: decimal.RequireFromString("100"),
Quantity: decimal.RequireFromString("2"),
Count: 1,
}},
Asks: []types.Level{{
Price: decimal.RequireFromString("101"),
Quantity: decimal.RequireFromString("3"),
Count: 1,
}},
}, 10)
if err != nil {
t.Fatalf("replace snapshot: %v", err)
}

nextBids := []*types.Level{
{
Price: decimal.RequireFromString("100"),
Quantity: decimal.RequireFromString("1.5"),
Count: 1,
},
}
nextAsks := []*types.Level{
{
Price: decimal.RequireFromString("102"),
Quantity: decimal.RequireFromString("4"),
Count: 1,
},
}
checksum := ComputeChecksum(nextBids, nextAsks)

err = book.ApplyDelta(Delta{
Symbol: "ETH-USDC",
Sequence: 11,
Bids: []types.Level{{
Price: decimal.RequireFromString("100"),
Quantity: decimal.RequireFromString("1.5"),
Count: 1,
}},
Asks: []types.Level{
{
Price: decimal.RequireFromString("101"),
Quantity: decimal.Zero,
Count: 0,
},
{
Price: decimal.RequireFromString("102"),
Quantity: decimal.RequireFromString("4"),
Count: 1,
},
},
Checksum: checksum,
})
if err != nil {
t.Fatalf("apply delta: %v", err)
}

snapshot := book.GetSnapshot()
if got := snapshot.Bids[0].Quantity.String(); got != "1.5" {
t.Fatalf("bid quantity = %s", got)
}
if got := snapshot.Asks[0].Price.String(); got != "102" {
t.Fatalf("ask price = %s", got)
}
if got := book.Sequence(); got != 11 {
t.Fatalf("sequence = %d", got)
}
}

func seededBook(t *testing.T) *OrderBook {
t.Helper()
book := NewOrderBook("BTC-USDC", Config{MaxDepth: 100})
err := book.ReplaceSnapshot(&types.DepthUpdate{
Symbol: "BTC-USDC",
Bids: []types.Level{{
Price: decimal.RequireFromString("100"),
Quantity: decimal.RequireFromString("2"),
Count: 1,
}},
Asks: []types.Level{{
Price: decimal.RequireFromString("101"),
Quantity: decimal.RequireFromString("3"),
Count: 1,
}},
}, 5)
if err != nil {
t.Fatalf("seed book: %v", err)
}
return book
}

func assertSnapshotEqual(t *testing.T, want *types.DepthUpdate, got *types.DepthUpdate) {
t.Helper()
if len(want.Bids) != len(got.Bids) || len(want.Asks) != len(got.Asks) {
t.Fatalf("snapshot size changed: want %+v got %+v", want, got)
}
for i := range want.Bids {
if !want.Bids[i].Price.Equal(got.Bids[i].Price) || !want.Bids[i].Quantity.Equal(got.Bids[i].Quantity) {
t.Fatalf("bid[%d] changed: want %+v got %+v", i, want.Bids[i], got.Bids[i])
}
}
for i := range want.Asks {
if !want.Asks[i].Price.Equal(got.Asks[i].Price) || !want.Asks[i].Quantity.Equal(got.Asks[i].Quantity) {
t.Fatalf("ask[%d] changed: want %+v got %+v", i, want.Asks[i], got.Asks[i])
}
}
}
Loading