implemented persistence
This commit is contained in:
@@ -0,0 +1,368 @@
|
||||
package internal_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nmoniz/any2anexoj/internal"
|
||||
"github.com/nmoniz/any2anexoj/internal/mocks"
|
||||
"github.com/shopspring/decimal"
|
||||
"go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
func TestFileStore_RoundTrip(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
ser := roundTripSerializer(ctrl)
|
||||
store, _ := newStore(t, "fake", ser)
|
||||
|
||||
original := map[string]*internal.FillerQueue{
|
||||
"AAA": newQueue(
|
||||
newFiller(ctrl, "AAA", 100, 50, 0),
|
||||
newFiller(ctrl, "AAA", 25, 80, 5),
|
||||
),
|
||||
"BBB": newQueue(
|
||||
newFiller(ctrl, "BBB", 7, 1000, 7),
|
||||
),
|
||||
}
|
||||
|
||||
if err := store.Save(t.Context(), original); err != nil {
|
||||
t.Fatalf("Save returned unexpected error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := store.Load(t.Context())
|
||||
if err != nil {
|
||||
t.Fatalf("Load returned unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if len(loaded) != len(original) {
|
||||
t.Fatalf("want %d symbols but got %d", len(original), len(loaded))
|
||||
}
|
||||
for symbol, wantQ := range original {
|
||||
gotQ, ok := loaded[symbol]
|
||||
if !ok {
|
||||
t.Fatalf("symbol %q missing from loaded state", symbol)
|
||||
}
|
||||
assertQueueEqual(t, symbol, gotQ, wantQ)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileStore_LoadMissingFileReturnsEmpty(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
ser := roundTripSerializer(ctrl)
|
||||
store, path := newStore(t, "fake", ser)
|
||||
|
||||
// Sanity: file doesn't exist.
|
||||
if _, err := os.Stat(path); !errors.Is(err, fs.ErrNotExist) {
|
||||
t.Fatalf("expected state file to be absent, got stat err: %v", err)
|
||||
}
|
||||
|
||||
queues, err := store.Load(t.Context())
|
||||
if err != nil {
|
||||
t.Fatalf("Load returned unexpected error for missing file: %v", err)
|
||||
}
|
||||
if queues == nil {
|
||||
t.Fatalf("Load returned nil map; want empty map")
|
||||
}
|
||||
if len(queues) != 0 {
|
||||
t.Fatalf("Load returned %d entries; want 0", len(queues))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileStore_LoadVersionMismatch(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
ser := roundTripSerializer(ctrl)
|
||||
store, path := newStore(t, "fake", ser)
|
||||
|
||||
// Hand-craft an unsupported-version state file.
|
||||
bad := struct {
|
||||
Version string `json:"version"`
|
||||
Platform string `json:"platform"`
|
||||
Queues map[string][]json.RawMessage `json:"queues"`
|
||||
}{
|
||||
Version: "999",
|
||||
Platform: "fake",
|
||||
Queues: map[string][]json.RawMessage{},
|
||||
}
|
||||
data, err := json.MarshalIndent(bad, "", " ")
|
||||
if err != nil {
|
||||
t.Fatalf("marshal: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatalf("write state: %v", err)
|
||||
}
|
||||
|
||||
_, err = store.Load(t.Context())
|
||||
if err == nil {
|
||||
t.Fatalf("Load with bad version should return an error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), `unexpected state version "999"`) {
|
||||
t.Errorf("expected version-mismatch error, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileStore_LoadPlatformMismatch(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
ser := roundTripSerializer(ctrl)
|
||||
store, path := newStore(t, "fake", ser)
|
||||
|
||||
// File claims a different platform.
|
||||
bad := struct {
|
||||
Version string `json:"version"`
|
||||
Platform string `json:"platform"`
|
||||
Queues map[string][]json.RawMessage `json:"queues"`
|
||||
}{
|
||||
Version: internal.StateVersion,
|
||||
Platform: "other-broker",
|
||||
Queues: map[string][]json.RawMessage{},
|
||||
}
|
||||
data, err := json.MarshalIndent(bad, "", " ")
|
||||
if err != nil {
|
||||
t.Fatalf("marshal: %v", err)
|
||||
}
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatalf("write state: %v", err)
|
||||
}
|
||||
|
||||
_, err = store.Load(t.Context())
|
||||
if err == nil {
|
||||
t.Fatalf("Load with mismatched platform should return an error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), `unexpected state platform "other-broker"`) {
|
||||
t.Errorf("expected platform-mismatch error, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileStore_SplitAdjustedLotSurvivesRoundTrip(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
ser := roundTripSerializer(ctrl)
|
||||
store, _ := newStore(t, "fake", ser)
|
||||
|
||||
// Simulate a lot that has been through a 5:1 split and is partially
|
||||
// filled. Starting from 10 shares @ $100, after a 5:1 split we should
|
||||
// have 50 shares @ $20 with 20 already filled.
|
||||
f := newFiller(ctrl, "SPLIT", 10, 100, 0)
|
||||
f.ApplySplit(decimal.NewFromInt(5))
|
||||
f.Fill(decimal.NewFromInt(20))
|
||||
|
||||
q := newQueue(f)
|
||||
|
||||
if err := store.Save(t.Context(), map[string]*internal.FillerQueue{"SPLIT": q}); err != nil {
|
||||
t.Fatalf("Save returned unexpected error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := store.Load(t.Context())
|
||||
if err != nil {
|
||||
t.Fatalf("Load returned unexpected error: %v", err)
|
||||
}
|
||||
gotQ := loaded["SPLIT"]
|
||||
if gotQ == nil || gotQ.Len() != 1 {
|
||||
t.Fatalf("want 1 lot for SPLIT, got %d", gotQ.Len())
|
||||
}
|
||||
got, _ := gotQ.Pop()
|
||||
if !got.Quantity().Equal(decimal.NewFromInt(50)) {
|
||||
t.Errorf("want quantity 50 but got %v", got.Quantity())
|
||||
}
|
||||
if !got.Price().Equal(decimal.NewFromInt(20)) {
|
||||
t.Errorf("want price 20 but got %v", got.Price())
|
||||
}
|
||||
if !got.Filled().Equal(decimal.NewFromInt(20)) {
|
||||
t.Errorf("want filled 20 but got %v", got.Filled())
|
||||
}
|
||||
if got.IsFilled() {
|
||||
t.Errorf("want IsFilled() to be false after split-adjusted partial fill")
|
||||
}
|
||||
|
||||
// Cost basis must round-trip exactly.
|
||||
if !got.Quantity().Mul(got.Price()).Equal(decimal.NewFromInt(1000)) {
|
||||
t.Errorf("want cost basis 1000 but got %v", got.Quantity().Mul(got.Price()))
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileStore_PartiallyFilledLotSurvivesRoundTrip(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
ser := roundTripSerializer(ctrl)
|
||||
store, _ := newStore(t, "fake", ser)
|
||||
|
||||
// A non-split lot that has been partially filled.
|
||||
f := newFiller(ctrl, "PART", 100, 50, 30)
|
||||
q := newQueue(f)
|
||||
|
||||
if err := store.Save(t.Context(), map[string]*internal.FillerQueue{"PART": q}); err != nil {
|
||||
t.Fatalf("Save returned unexpected error: %v", err)
|
||||
}
|
||||
|
||||
loaded, err := store.Load(t.Context())
|
||||
if err != nil {
|
||||
t.Fatalf("Load returned unexpected error: %v", err)
|
||||
}
|
||||
gotQ := loaded["PART"]
|
||||
if gotQ == nil || gotQ.Len() != 1 {
|
||||
t.Fatalf("want 1 lot for PART, got %d", gotQ.Len())
|
||||
}
|
||||
got, _ := gotQ.Pop()
|
||||
if !got.Quantity().Equal(decimal.NewFromInt(100)) {
|
||||
t.Errorf("want quantity 100 but got %v", got.Quantity())
|
||||
}
|
||||
if !got.Price().Equal(decimal.NewFromInt(50)) {
|
||||
t.Errorf("want price 50 but got %v", got.Price())
|
||||
}
|
||||
if !got.Filled().Equal(decimal.NewFromInt(30)) {
|
||||
t.Errorf("want filled 30 but got %v", got.Filled())
|
||||
}
|
||||
if got.IsFilled() {
|
||||
t.Errorf("want IsFilled() to be false (30/100 filled)")
|
||||
}
|
||||
// Filling the remaining 70 must now make it filled.
|
||||
_, done := got.Fill(decimal.NewFromInt(70))
|
||||
if !done {
|
||||
t.Errorf("after filling remaining 70, IsFilled() should be true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewFileStore_ValidatesArguments(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
ser := mocks.NewMockRecordSerializer(ctrl)
|
||||
|
||||
if _, err := internal.NewFileStore("", "fake", ser); err == nil {
|
||||
t.Errorf("NewFileStore with empty filename should fail")
|
||||
}
|
||||
if _, err := internal.NewFileStore("/tmp/x", "fake", nil); err == nil {
|
||||
t.Errorf("NewFileStore with nil serializer should fail")
|
||||
}
|
||||
}
|
||||
|
||||
// newRecord builds a MockRecord whose Symbol() returns the given symbol.
|
||||
// The FileStore only reads Symbol() off the loaded Record during tests
|
||||
// (quantity/price/filled come from the persisted struct fields), so all
|
||||
// other Record methods can be left as default-mocked values.
|
||||
func newRecord(ctrl *gomock.Controller, symbol string) *mocks.MockRecord {
|
||||
r := mocks.NewMockRecord(ctrl)
|
||||
r.EXPECT().Symbol().Return(symbol).AnyTimes()
|
||||
return r
|
||||
}
|
||||
|
||||
// newFiller builds a Filler backed by a MockRecord with the given symbol,
|
||||
// quantity, price, and filled amounts (in whole units). It mirrors the
|
||||
// inline calls that previously littered every test.
|
||||
func newFiller(ctrl *gomock.Controller, symbol string, quantity, price, filled int64) *internal.Filler {
|
||||
return internal.NewFillerFromState(
|
||||
newRecord(ctrl, symbol),
|
||||
decimal.NewFromInt(quantity),
|
||||
decimal.NewFromInt(price),
|
||||
decimal.NewFromInt(filled),
|
||||
)
|
||||
}
|
||||
|
||||
// newQueue creates a FillerQueue pre-populated with the given fillers.
|
||||
func newQueue(fillers ...*internal.Filler) *internal.FillerQueue {
|
||||
q := new(internal.FillerQueue)
|
||||
for _, f := range fillers {
|
||||
q.Push(f)
|
||||
}
|
||||
return q
|
||||
}
|
||||
|
||||
// newStore creates a FileStore backed by a temp file and returns the store
|
||||
// plus the path to the underlying state file.
|
||||
func newStore(t *testing.T, platform string, ser internal.RecordSerializer) (*internal.FileStore, string) {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "state.json")
|
||||
store, err := internal.NewFileStore(path, platform, ser)
|
||||
if err != nil {
|
||||
t.Fatalf("NewFileStore returned unexpected error: %v", err)
|
||||
}
|
||||
return store, path
|
||||
}
|
||||
|
||||
// roundTripSerializer returns a serializer mock whose MarshalRecord encodes
|
||||
// the Symbol into bytes and whose UnmarshalRecord decodes those bytes back
|
||||
// into a fresh MockRecord returning the encoded Symbol. Tests use this when
|
||||
// they need Save + Load to round-trip equivalent Records.
|
||||
func roundTripSerializer(ctrl *gomock.Controller) *mocks.MockRecordSerializer {
|
||||
ser := mocks.NewMockRecordSerializer(ctrl)
|
||||
ser.EXPECT().
|
||||
MarshalRecord(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(_ context.Context, r internal.Record) ([]byte, error) {
|
||||
return []byte(r.Symbol()), nil
|
||||
}).
|
||||
AnyTimes()
|
||||
ser.EXPECT().
|
||||
UnmarshalRecord(gomock.Any(), gomock.Any()).
|
||||
DoAndReturn(func(_ context.Context, b []byte) (internal.Record, error) {
|
||||
return newRecord(ctrl, string(b)), nil
|
||||
}).
|
||||
AnyTimes()
|
||||
return ser
|
||||
}
|
||||
|
||||
// queueSnapshot drains the given FillerQueue (via Pop) and returns its
|
||||
// contents in a side-effect-free shape suitable for value comparison.
|
||||
// Callers should not use the queue afterwards.
|
||||
func queueSnapshot(q *internal.FillerQueue) []queueEntry {
|
||||
if q == nil {
|
||||
return nil
|
||||
}
|
||||
var out []queueEntry
|
||||
for {
|
||||
f, ok := q.Pop()
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
out = append(out, queueEntry{
|
||||
symbol: f.Symbol(),
|
||||
quantity: f.Quantity(),
|
||||
price: f.Price(),
|
||||
filled: f.Filled(),
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// assertQueueEqual compares two FillerQueues by popping every element from
|
||||
// each and comparing the resulting sequence of queueEntries. Both queues are
|
||||
// drained as a side effect.
|
||||
func assertQueueEqual(t *testing.T, symbol string, got, want *internal.FillerQueue) {
|
||||
t.Helper()
|
||||
if got == nil {
|
||||
t.Fatalf("symbol %q: loaded queue is nil", symbol)
|
||||
}
|
||||
if got.Len() != want.Len() {
|
||||
t.Fatalf("symbol %q: want %d lots but got %d", symbol, want.Len(), got.Len())
|
||||
}
|
||||
wantEntries := queueSnapshot(want)
|
||||
gotEntries := queueSnapshot(got)
|
||||
for i := range wantEntries {
|
||||
w := wantEntries[i]
|
||||
g := gotEntries[i]
|
||||
if w.symbol != g.symbol {
|
||||
t.Errorf("symbol %q lot %d: want symbol %q but got %q",
|
||||
symbol, i, w.symbol, g.symbol)
|
||||
}
|
||||
if !w.quantity.Equal(g.quantity) {
|
||||
t.Errorf("symbol %q lot %d: want quantity %v but got %v",
|
||||
symbol, i, w.quantity, g.quantity)
|
||||
}
|
||||
if !w.price.Equal(g.price) {
|
||||
t.Errorf("symbol %q lot %d: want price %v but got %v",
|
||||
symbol, i, w.price, g.price)
|
||||
}
|
||||
if !w.filled.Equal(g.filled) {
|
||||
t.Errorf("symbol %q lot %d: want filled %v but got %v",
|
||||
symbol, i, w.filled, g.filled)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type queueEntry struct {
|
||||
symbol string
|
||||
quantity decimal.Decimal
|
||||
price decimal.Decimal
|
||||
filled decimal.Decimal
|
||||
}
|
||||
Reference in New Issue
Block a user