diff --git a/market/matching/delta_validation_test.go b/market/matching/delta_validation_test.go new file mode 100644 index 00000000..02b21bb4 --- /dev/null +++ b/market/matching/delta_validation_test.go @@ -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) + } +} diff --git a/market/orderbook/delta.go b/market/orderbook/delta.go new file mode 100644 index 00000000..7a5428f0 --- /dev/null +++ b/market/orderbook/delta.go @@ -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, ©) + } + 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, ©) + } + return result +} diff --git a/market/orderbook/delta_test.go b/market/orderbook/delta_test.go new file mode 100644 index 00000000..f86b67f3 --- /dev/null +++ b/market/orderbook/delta_test.go @@ -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]) + } + } +} diff --git a/market/orderbook/orderbook.go b/market/orderbook/orderbook.go index 98f0bc5b..963effc8 100644 --- a/market/orderbook/orderbook.go +++ b/market/orderbook/orderbook.go @@ -169,13 +169,17 @@ func (e *BookError) Error() string { func insertLevel(levels []*types.Level, level *types.Level, desc bool) []*types.Level { levels = append(levels, level) + sortLevels(levels, desc) + return levels +} + +func sortLevels(levels []*types.Level, desc bool) { sort.Slice(levels, func(i, j int) bool { if desc { return levels[i].Price.GreaterThan(levels[j].Price) } return levels[i].Price.LessThan(levels[j].Price) }) - return levels } func removeLevel(levels []*types.Level, price decimal.Decimal) []*types.Level { diff --git a/market/ws/delta.go b/market/ws/delta.go new file mode 100644 index 00000000..f505b7ea --- /dev/null +++ b/market/ws/delta.go @@ -0,0 +1,33 @@ +package ws + +import ( + "encoding/json" + "errors" + "fmt" + + "github.com/tent-of-trials/market/orderbook" +) + +var ErrInvalidDeltaMessage = errors.New("invalid order book delta message") + +type deltaEnvelope struct { + Type string `json:"type"` + Delta orderbook.Delta `json:"delta"` +} + +func ParseOrderBookDeltaMessage(message []byte) (orderbook.Delta, error) { + var envelope deltaEnvelope + if err := json.Unmarshal(message, &envelope); err != nil { + return orderbook.Delta{}, fmt.Errorf("%w: %v", ErrInvalidDeltaMessage, err) + } + if envelope.Type != "orderbook_delta" { + return orderbook.Delta{}, fmt.Errorf("%w: unsupported type %q", ErrInvalidDeltaMessage, envelope.Type) + } + if envelope.Delta.Symbol == "" { + return orderbook.Delta{}, fmt.Errorf("%w: missing symbol", ErrInvalidDeltaMessage) + } + if envelope.Delta.Sequence == 0 { + return orderbook.Delta{}, fmt.Errorf("%w: missing sequence", ErrInvalidDeltaMessage) + } + return envelope.Delta, nil +} diff --git a/market/ws/delta_test.go b/market/ws/delta_test.go new file mode 100644 index 00000000..8e9cb08d --- /dev/null +++ b/market/ws/delta_test.go @@ -0,0 +1,37 @@ +package ws + +import ( + "errors" + "testing" +) + +func TestParseOrderBookDeltaMessageRejectsMalformedPayload(t *testing.T) { + _, err := ParseOrderBookDeltaMessage([]byte(`{"type":"orderbook_delta","delta":{"symbol":"BTC-USDC","sequence":1,"bids":[{"price":]}}`)) + if !errors.Is(err, ErrInvalidDeltaMessage) { + t.Fatalf("expected invalid delta message error, got %v", err) + } +} + +func TestParseOrderBookDeltaMessageRejectsWrongEventType(t *testing.T) { + _, err := ParseOrderBookDeltaMessage([]byte(`{"type":"trade","delta":{"symbol":"BTC-USDC","sequence":1}}`)) + if !errors.Is(err, ErrInvalidDeltaMessage) { + t.Fatalf("expected invalid delta message error, got %v", err) + } +} + +func TestParseOrderBookDeltaMessageAcceptsValidDelta(t *testing.T) { + delta, err := ParseOrderBookDeltaMessage([]byte(`{ + "type":"orderbook_delta", + "delta":{ + "symbol":"BTC-USDC", + "sequence":2, + "bids":[{"price":"100","quantity":"1","order_count":1}] + } + }`)) + if err != nil { + t.Fatalf("parse valid delta: %v", err) + } + if delta.Symbol != "BTC-USDC" || delta.Sequence != 2 || len(delta.Bids) != 1 { + t.Fatalf("unexpected delta: %+v", delta) + } +}