import unittest

from strategy_engine import StrategyEngine


class FakeTradeApi:
    def __init__(self, fills):
        self.fills = fills

    def get_user_trades_by_instrument_and_time(
        self, instrument_name, start_timestamp, end_timestamp, count=1000
    ):
        return {
            "success": True,
            "result": {"trades": list(self.fills), "has_more": False},
        }


class TradeReconciliationTests(unittest.TestCase):
    def make_engine(self, fills):
        engine = StrategyEngine("id", "secret", testnet=True)
        engine.api = FakeTradeApi(fills)
        engine.trades = []
        engine.total_trades = 0
        engine.strategy_started_at_ms = 0
        engine._save_state = lambda: None
        return engine

    @staticmethod
    def fill(order_id, trade_id, side, amount, price, timestamp):
        return {
            "instrument_name": "BTC_USDC",
            "label": f"maker_{side}",
            "order_id": order_id,
            "trade_id": trade_id,
            "direction": side,
            "amount": amount,
            "price": price,
            "timestamp": timestamp,
        }

    def test_partial_cancelled_order_is_aggregated_and_idempotent(self):
        fills = [
            self.fill("sell-1", "fill-1", "sell", 0.0001, 65000, 1784800000000),
            self.fill("sell-1", "fill-2", "sell", 0.0002, 65010, 1784800005000),
        ]
        engine = self.make_engine(fills)

        self.assertTrue(engine._reconcile_exchange_trades(force=True))
        self.assertEqual(engine.total_trades, 1)
        self.assertEqual(engine.trades[0]["amount_btc"], 0.0003)
        self.assertEqual(engine.trades[0]["total_usdc"], 19.5)
        self.assertTrue(engine.trades[0]["recovered"])

        self.assertFalse(engine._reconcile_exchange_trades(force=True))
        self.assertEqual(engine.total_trades, 1)

    def test_tracked_order_sync_is_not_marked_historical_recovery(self):
        fills = [
            self.fill("buy-1", "fill-1", "buy", 0.0015, 64715, 1784800100000),
        ]
        engine = self.make_engine(fills)
        engine._our_buy_id = "buy-1"

        engine._reconcile_exchange_trades(force=True)

        self.assertFalse(engine.trades[0]["recovered"])
        self.assertEqual(engine.trades[0]["source"], "exchange_sync")

    def test_forward_pnl_uses_buy_and_hold_baseline(self):
        engine = self.make_engine([])
        engine.initial_usdc = 100000
        engine.initial_btc = 100
        engine.initial_total_usdc = 6600000
        engine.btc_index_price = 65000
        engine.usdc_balance = 100164.5
        engine.btc_balance = 99.9975
        engine.btc_value_usdc = engine.btc_balance * engine.btc_index_price
        engine.total_value_usdc = engine.usdc_balance + engine.btc_value_usdc

        state = engine.get_state()

        self.assertEqual(state["buy_and_hold_value"], 6600000)
        self.assertAlmostEqual(state["forward_pnl"], 2.0)
        self.assertAlmostEqual(state["total_pnl"], 2.0)
        self.assertAlmostEqual(state["forward_return_pct"], 0.002)


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