"""Synthetic inputs for the upgrade guide; backtests require normal FinLab login.

Run outside a source checkout: python upgrade_results.py CASE
Quarterly/calendar fixtures replace only data inputs, never the calculation.
"""

from __future__ import annotations

import sys
from unittest.mock import patch

import numpy as np
import pandas as pd

import finlab
from finlab import data
from finlab.backtest import sim
from finlab.dataframe import FinlabDataFrame as F
from finlab.market import Market
from finlab.tools.event_study import create_factor_data
from finlab.tools.factor_metrics import calc_ic


class Prices(Market):
    def __init__(self, close: pd.DataFrame) -> None:
        self.close = close

    @staticmethod
    def get_name() -> str:
        return "synthetic"

    def get_asset_id_to_name(self) -> dict:
        return {}

    def get_price(self, trade_at_price: object, adj: bool = True) -> pd.DataFrame:
        return self.close.copy()


def backtest(position: pd.DataFrame, close: pd.DataFrame, **kwargs: object) -> object:
    return sim(
        position, market=Prices(close), upload=False, fee_ratio=0, tax_ratio=0, **kwargs
    )


def quarterly(case: str) -> object:
    quarters = ["2023-Q1", "2023-Q2"]
    keys = pd.DataFrame(
        {
            "A": pd.to_datetime(["2023-05-01", "2023-08-01"]),
            "B": pd.to_datetime(["2023-05-02", "2023-08-02"]),
        },
        index=quarters,
    )
    q = F({"A": [5.0, 7.0], "B": [9.0, 3.0]}, index=quarters)
    if case == "quarter_bool":
        q = q > 6
    if case == "latest_quarter":
        keys.iloc[:, :] = pd.Timestamp("2023-08-01")
    if case == "missing_date":
        keys.loc["2023-Q1", "A"] = pd.NaT
    with patch.object(
        data,
        "get",
        side_effect=lambda name: (
            F(
                100.0,
                index=pd.bdate_range("2023-05-01", "2023-09-01"),
                columns=["A", "B"],
            )
            if name == "price:收盤價"
            else keys
        ),
    ):
        if case == "quarter_bool":
            return (~q.index_str_to_date()).loc["2023-05-02"].tolist()
        if case == "latest_quarter":
            return q.index_str_to_date().iloc[-1].tolist()
        if case == "quarter_reindex":
            return q.reindex(pd.to_datetime(["2023-05-01"])).iloc[0].tolist()
        if case == "quarter_rank":
            out = q.rank(axis=1, pct=True).index_str_to_date()
            return out.loc["2023-05-01", "A"]
        if case == "missing_date":
            return q.deadline().loc["2023-05-02", "A"]
        if case == "quarter_group":
            with patch(
                "finlab.dataframe.core._fetch_refined_categories",
                return_value={"A": "x", "B": "x"},
            ):
                out = q.groupby_category().mean()
                return [str(out.index[0]), out.iloc[0, 0]]

    raise ValueError(case)


def clock_case() -> object:
    import datetime as dt
    import importlib
    from types import SimpleNamespace

    from finlab.markets.tw import TWMarket

    instant = dt.datetime.fromisoformat("2024-08-30T16:30:00+00:00")
    original = dt.datetime

    class Frozen(original):
        @classmethod
        def now(cls, tz: object = None) -> object:
            return (
                instant.astimezone(tz)
                if tz
                else instant.astimezone().replace(tzinfo=None)
            )

    dates = pd.bdate_range("2024-01-01", "2024-08-30")
    prices = F(100.0, index=dates, columns=["2330", "2317"])

    class Taiwan(TWMarket):
        def get_price(self, trade_at_price: object, adj: bool = True) -> pd.DataFrame:
            return prices.copy()

        def get_asset_id_to_name(self) -> dict:
            return {}

        def get_industry(self) -> dict:
            return {}

        def market_close_at_timestamp(self, timestamp: object = None) -> object:
            return super().market_close_at_timestamp(
                instant if timestamp is None else timestamp
            )

        def get_delisting_notices(self) -> pd.DataFrame:
            return pd.DataFrame(columns=["announced", "delisted"])

    pos = F(
        {"2330": [True, False, True, False], "2317": [False, True, False, True]},
        index=pd.to_datetime(["2024-08-28", "2024-08-29", "2024-08-30", "2024-08-31"]),
    )
    module = importlib.import_module("finlab.backtest.sim")
    with patch.object(module.datetime, "datetime", Frozen):
        try:
            calendar = importlib.import_module("finlab.data.calendar")
        except ImportError:
            calendar = None
        fixture = (
            patch.object(
                calendar,
                "get_calendar",
                return_value=SimpleNamespace(is_session=lambda day: day.weekday() < 5),
            )
            if calendar
            else patch.object(data, "get", return_value=prices)
        )
        with fixture:
            report = sim(pos, market=Taiwan(), resample=None, upload=False)
            return report.next_weights[report.next_weights != 0].to_dict()


def calculation(case: str) -> object:  # noqa: C901 - independent documented cases
    index = pd.bdate_range("2022-01-03", periods=80)
    close = F(
        100 * 1.001 ** np.arange(80)[:, None] * np.ones((80, 2)),
        index=index,
        columns=["A", "B"],
    )
    pos = F(True, index=index, columns=close.columns)
    if case in {"universe_reset", "universe_empty"}:
        import importlib

        module = importlib.import_module("finlab.data.universe")
        categories = pd.DataFrame(
            {
                "stock_id": ["2330", "2317"],
                "market": ["sii", "sii"],
                "category": ["半導體業", "其他電子業"],
            }
        )
        with (
            patch.object(data, "get", return_value=categories),
            patch.object(module, "universe_stocks", set()),
        ):
            data.set_universe(
                **({"sector": "打錯的產業"} if case == "universe_empty" else {})
            )
            return module.refine_stock_id(
                "price:收盤價",
                F(1.0, index=index[:1], columns=["2330", "2317", "0050O"]),
            ).columns.tolist()
    if case == "clock":
        return clock_case()
    if case == "from_weight":
        from finlab.online.core.position import Position

        return Position.from_weight(
            {"2330": 0.8, "2317": 0.8}, fund=1000000, price={"2330": 100, "2317": 100}
        ).to_list()
    if case == "pandas_left":
        left = pd.DataFrame({"A": [1.0]}, index=index[:1])
        right = F({"B": [2.0]}, index=index[1:2])
        out = left * right
        return [type(out).__name__, out.shape]
    if case == "series":
        return (close.iloc[:1] * pd.Series({"A": 2, "B": 3})).iloc[0].tolist()
    if case == "timing":
        return pos.iloc[:3][pd.Series([True, False, True], index=index[:3])][
            "A"
        ].tolist()
    if case == "hold_default":
        entries = F({"A": [True, True, True]}, index=index[:3])
        exits = F({"A": [False, True]}, index=index[:2])
        return entries.hold_until(exits)["A"].tolist()
    if case == "rank_ic":
        ix = pd.MultiIndex.from_product(
            [[index[0]], ["A", "B", "C", "D"]], names=["date", "stock"]
        )
        return calc_ic(
            pd.DataFrame({"f": [1, 2, 3, 4]}, index=ix),
            pd.Series([1, 2, 3, 100], index=ix),
            rank=True,
        ).iloc[0, 0]
    if case == "quantiles":
        prices = F(100.0, index=index, columns=[str(i) for i in range(10)])
        factor = F([list(range(10))], index=index[:1], columns=prices.columns)
        with patch.object(data, "get", return_value=prices):
            return create_factor_data(factor, prices, days=[1])[
                "factor_factor_quantile"
            ].tolist()
    if case in ("drawdown", "volatility"):
        close.iloc[1:] = close.iloc[1:].mul(
            (1 + 0.03 * np.sin(np.arange(1, 80))) * 0.95 ** np.arange(1, 80), axis=0
        )
        weights = (
            pos.weight.drawdown_control(0.1, price=close)
            if case == "drawdown"
            else pos.weight.target_volatility(0.01, window=3, price=close)
        )
        return weights.abs().sum(axis=1).iloc[[0, 5]].tolist()
    if case == "nan_position":
        pos = F({"A": [True, np.nan, False, True], "B": [False] * 4}, index=index[:4])
        out = backtest(pos, close, resample=None)
        return out.trades["entry_date"].dt.strftime("%Y-%m-%d").tolist()
    if case == "short_position":
        return backtest(pos.iloc[:2], close).creturn.index[-1].strftime("%Y-%m-%d")
    if case == "future_position":
        pos = pos.iloc[[0, 5, 10]].copy()
        pos.loc[pd.Timestamp("2030-01-01")] = False
        return (
            backtest(pos, close, resample=None).creturn.index[-1].strftime("%Y-%m-%d")
        )
    if case == "delisted":
        close.loc[index[25] :, "B"] = np.nan
        out = backtest(pos, close, resample=None)
        return out.position.reindex(index, method="ffill").loc[index[40], "B"]
    if case in {"stop_target", "stop_reentry"}:
        dates = pd.bdate_range("2022-01-03", "2022-04-15")
        symbols = ["2330", "2317", "2454", "2603", "1301"]
        prices = F(100.0, index=dates, columns=symbols)
        entries = F(True, index=dates, columns=symbols)
        if case == "stop_target":
            prices.loc["2022-03-15":, ["2330", "2454", "2603"]] *= 0.9
            out = backtest(
                entries.loc[:"2022-03-31"],
                prices.loc[:"2022-03-31"],
                resample="M",
                stop_loss=0.05,
            )
        else:
            prices.loc["2022-03-15":, "2330"] *= 0.9
            entries.loc[:"2022-03-30", "2454"] = False
            entries.loc[:, ["2603", "1301"]] = False
            out = backtest(entries, prices, resample=None, stop_loss=0.05)
        return out.next_weights[out.next_weights != 0].to_dict()
    if case == "monthly_saturday":
        prices = F(
            100.0,
            index=pd.to_datetime(["2016-09-09", "2016-09-10", "2016-09-12"]),
            columns=["A"],
        )
        with patch.object(data, "get", return_value=prices):
            return str(
                F({"A": [1]}, index=pd.to_datetime(["2016-09-10"]))
                ._index_to_business_day()
                .index[0]
                .date()
            )
    if case == "factor_saturday":
        prices = F(
            100.0,
            index=pd.to_datetime(
                ["2016-09-09", "2016-09-10", "2016-09-12", "2016-09-13", "2016-09-14"]
            ),
            columns=["A"],
        )
        factor = F({"A": [1]}, index=pd.to_datetime(["2016-09-10"]))
        with patch.object(data, "get", return_value=prices):
            return str(create_factor_data(factor, prices, days=[1]).index[0][0].date())
    if case == "indicator_typo":
        with patch.object(data, "get", return_value=close):
            return data.indicator("SMA", timperiod=3).iloc[5, 0]
    if case == "duplicate_columns":
        from finlab.ffn_core import drop_duplicate_cols

        return drop_duplicate_cols(
            pd.DataFrame([[1, 2, 3, 4, 5, 6]], columns=["a", "a", "b", "b", "c", "c"])
        ).columns.tolist()
    raise ValueError(case)


if __name__ == "__main__":
    print(
        "finlab", finlab.__version__, "pandas", pd.__version__, "numpy", np.__version__
    )
    case = sys.argv[1]
    # Authenticate normally for the cases calling sim(). See the login guide.
    if case in {
        "nan_position",
        "short_position",
        "future_position",
        "delisted",
        "stop_target",
        "stop_reentry",
        "clock",
    }:
        finlab.login()
    try:
        value = (
            quarterly(case)
            if case.startswith("quarter_") or case in {"latest_quarter", "missing_date"}
            else calculation(case)
        )
        print(case, value)
    except (ValueError, TypeError, IndexError, AssertionError, KeyError) as error:
        print(case, type(error).__name__, str(error).splitlines()[0])
