import unittest

from paper_trader import PaperTrader


class PaperTraderTests(unittest.TestCase):
    def test_starts_with_virtual_cash_and_no_positions(self):
        trader = PaperTrader(cash=10_000, fast_window=2, slow_window=3)
        state = trader.snapshot({"BTC/USD": 100})
        self.assertEqual(state["cash"], 10_000)
        self.assertEqual(state["positions"], {})
        self.assertEqual(state["equity"], 10_000)

    def test_buys_after_fast_average_crosses_above_slow_average(self):
        trader = PaperTrader(cash=10_000, fast_window=2, slow_window=3, allocation=0.5)
        for price in [100, 100, 110, 120]:
            trader.step("BTC/USD", price)
        state = trader.snapshot({"BTC/USD": 120})
        self.assertGreater(state["positions"]["BTC/USD"]["quantity"], 0)
        self.assertEqual(state["trades"][0]["side"], "buy")
        self.assertLess(state["cash"], 10_000)

    def test_sells_position_after_fast_average_crosses_below_slow_average(self):
        trader = PaperTrader(cash=10_000, fast_window=2, slow_window=3, allocation=0.5)
        for price in [100, 100, 110, 120, 80, 70]:
            trader.step("BTC/USD", price)
        state = trader.snapshot({"BTC/USD": 70})
        self.assertEqual(state["positions"], {})
        self.assertEqual(state["trades"][-1]["side"], "sell")
        self.assertLess(state["cash"], 10_000)


if __name__ == "__main__":
    unittest.main()
