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 }