#!/usr/bin/env python3
"""Reproduce the Day 8 Bitcoin price, volatility, and drawdown analysis.

The default run uses a fixed Coin Metrics PriceUSD file covering 2011-08-27 to
2026-08-26. Use --refresh to download that date window again before running the
analysis.
"""

from __future__ import annotations

import argparse
import math
from pathlib import Path
from urllib.parse import urlencode
from urllib.request import urlopen

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import seaborn as sns


ROOT = Path(__file__).resolve().parents[1]
DEFAULT_RAW = ROOT / "data" / "raw" / "btc_priceusd_2011_2026.csv"
DEFAULT_OUTPUT = ROOT / "data" / "analysis"
DEFAULT_FIGURES = ROOT / "figures"

API = "https://community-api.coinmetrics.io/v4/timeseries/asset-metrics"
START_DATE = "2011-08-27"
END_DATE = "2026-08-26"
TRADING_DAYS = 365  # Bitcoin trades every calendar day.

NAVY = "#002147"
CRIMSON = "#AA381E"
GOLD = "#D4AF37"
TEAL = "#2A7F86"
GRAY = "#746F67"
PALE = "#F6F8FA"


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--input", type=Path, default=DEFAULT_RAW)
    parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
    parser.add_argument("--figure-dir", type=Path, default=DEFAULT_FIGURES)
    parser.add_argument(
        "--refresh",
        action="store_true",
        help="Download the fixed course window from Coin Metrics before analysis.",
    )
    return parser.parse_args()


def download_price_history(destination: Path) -> None:
    """Download the fixed 15-year daily PriceUSD window without an API key."""
    params = {
        "assets": "btc",
        "metrics": "PriceUSD",
        "frequency": "1d",
        "start_time": START_DATE,
        "end_time": END_DATE,
        "page_size": 10000,
        "paging_from": "start",
        "format": "csv",
    }
    url = f"{API}?{urlencode(params)}"
    destination.parent.mkdir(parents=True, exist_ok=True)
    with urlopen(url, timeout=60) as response:
        payload = response.read()
    if not payload.startswith(b"asset,time,PriceUSD"):
        raise RuntimeError("Coin Metrics did not return the expected PriceUSD CSV.")
    destination.write_bytes(payload)


def load_prices(raw_path: Path) -> pd.DataFrame:
    raw = pd.read_csv(raw_path)
    required = {"asset", "time", "PriceUSD"}
    if not required.issubset(raw.columns):
        raise ValueError(f"Expected columns {sorted(required)}; got {raw.columns.tolist()}")

    data = raw.rename(columns={"time": "date", "PriceUSD": "price_usd"}).copy()
    data["date"] = pd.to_datetime(data["date"], utc=True).dt.tz_localize(None)
    data["price_usd"] = pd.to_numeric(data["price_usd"], errors="raise")
    data = data[["date", "price_usd"]].sort_values("date").drop_duplicates("date")
    data = data.reset_index(drop=True)

    if data.empty or not data["price_usd"].gt(0).all():
        raise ValueError("Bitcoin prices must be present and strictly positive.")
    if not data["date"].is_monotonic_increasing or not data["date"].is_unique:
        raise ValueError("Dates must be unique and increasing.")
    span_years = (data["date"].iloc[-1] - data["date"].iloc[0]).days / 365.25
    if span_years < 14.9:
        raise ValueError(f"Expected about 15 years of data; found {span_years:.2f} years.")
    return data


def calculate_metrics(data: pd.DataFrame) -> pd.DataFrame:
    out = data.copy()
    out["simple_return"] = out["price_usd"].pct_change()
    out["log_return"] = np.log(out["price_usd"]).diff()
    out["volatility_30d_ann"] = out["log_return"].rolling(30).std() * math.sqrt(TRADING_DAYS)
    out["volatility_365d_ann"] = out["log_return"].rolling(365).std() * math.sqrt(TRADING_DAYS)
    out["running_peak_usd"] = out["price_usd"].cummax()
    out["drawdown"] = out["price_usd"] / out["running_peak_usd"] - 1
    out["recovery_needed"] = 1 / (1 + out["drawdown"]) - 1
    return out


def headline_summary(data: pd.DataFrame) -> pd.DataFrame:
    returns = data["log_return"].dropna()
    first = data.iloc[0]
    last = data.iloc[-1]
    years = (last["date"] - first["date"]).days / 365.25
    trough_i = data["drawdown"].idxmin()
    best_i = data["simple_return"].idxmax()
    worst_i = data["simple_return"].idxmin()
    values = {
        "sample_start": first["date"].date().isoformat(),
        "sample_end": last["date"].date().isoformat(),
        "observations": int(len(data)),
        "start_price_usd": float(first["price_usd"]),
        "end_price_usd": float(last["price_usd"]),
        "compound_annual_growth_rate": float(
            (last["price_usd"] / first["price_usd"]) ** (1 / years) - 1
        ),
        "full_sample_annualised_volatility": float(returns.std() * math.sqrt(TRADING_DAYS)),
        "maximum_drawdown": float(data.loc[trough_i, "drawdown"]),
        "maximum_drawdown_date": data.loc[trough_i, "date"].date().isoformat(),
        "best_daily_return": float(data.loc[best_i, "simple_return"]),
        "best_daily_return_date": data.loc[best_i, "date"].date().isoformat(),
        "worst_daily_return": float(data.loc[worst_i, "simple_return"]),
        "worst_daily_return_date": data.loc[worst_i, "date"].date().isoformat(),
        "latest_30d_annualised_volatility": float(data["volatility_30d_ann"].dropna().iloc[-1]),
        "latest_365d_annualised_volatility": float(data["volatility_365d_ann"].dropna().iloc[-1]),
    }
    return pd.DataFrame({"metric": values.keys(), "value": values.values()})


def annual_summary(data: pd.DataFrame) -> pd.DataFrame:
    annual = []
    for year, group in data.groupby(data["date"].dt.year):
        rets = group["log_return"].dropna()
        if len(group) < 2 or rets.empty:
            continue
        complete_year = (
            group["date"].min() == pd.Timestamp(year=int(year), month=1, day=1)
            and group["date"].max() == pd.Timestamp(year=int(year), month=12, day=31)
        )
        running_peak = group["price_usd"].cummax()
        annual.append(
            {
                "year": int(year),
                "observations": int(len(group)),
                "complete_year": bool(complete_year),
                "annual_return": float(group["price_usd"].iloc[-1] / group["price_usd"].iloc[0] - 1),
                "annualised_volatility": float(rets.std() * math.sqrt(TRADING_DAYS)),
                "maximum_within_year_drawdown": float((group["price_usd"] / running_peak - 1).min()),
            }
        )
    return pd.DataFrame(annual)


def compare_volatility_windows(data: pd.DataFrame) -> pd.Series:
    """Compare the first and most recent five complete calendar years."""
    early = data.loc[data["date"].dt.year.between(2012, 2016), "log_return"]
    recent = data.loc[data["date"].dt.year.between(2021, 2025), "log_return"]
    return pd.Series(
        {
            "2012-2016": early.dropna().std() * math.sqrt(TRADING_DAYS),
            "2021-2025": recent.dropna().std() * math.sqrt(TRADING_DAYS),
        },
        name="annualised_volatility",
    )


def configure_style() -> None:
    sns.set_theme(style="whitegrid", context="talk")
    plt.rcParams.update(
        {
            "font.family": "DejaVu Sans",
            "axes.titleweight": "bold",
            "axes.titlecolor": NAVY,
            "axes.labelcolor": NAVY,
            "axes.edgecolor": "#CFC8BD",
            "grid.color": "#E4DED5",
            "figure.facecolor": "white",
            "axes.facecolor": "white",
        }
    )


def plot_history(data: pd.DataFrame, figure_dir: Path) -> None:
    configure_style()
    figure_dir.mkdir(parents=True, exist_ok=True)

    fig, axes = plt.subplots(3, 1, figsize=(13, 11), sharex=True, constrained_layout=True)

    axes[0].plot(data["date"], data["price_usd"], color=NAVY, linewidth=1.5)
    axes[0].set_yscale("log")
    axes[0].set_ylabel("USD, log scale")
    axes[0].set_title("Bitcoin price: a long sample contains several distinct regimes", loc="left")

    axes[1].plot(
        data["date"], 100 * data["volatility_30d_ann"],
        color=CRIMSON, linewidth=1.0, alpha=0.8, label="30-day",
    )
    axes[1].plot(
        data["date"], 100 * data["volatility_365d_ann"],
        color=TEAL, linewidth=1.8, label="365-day",
    )
    axes[1].set_ylabel("Annualised volatility (%)")
    axes[1].set_title("Volatility depends on the observation window", loc="left")
    axes[1].legend(frameon=False, ncol=2, loc="upper right")

    axes[2].fill_between(data["date"], 100 * data["drawdown"], 0, color=CRIMSON, alpha=0.28)
    axes[2].plot(data["date"], 100 * data["drawdown"], color=CRIMSON, linewidth=1.0)
    axes[2].axhline(0, color=GRAY, linewidth=0.8)
    axes[2].set_ylabel("Drawdown (%)")
    axes[2].set_xlabel("Date")
    axes[2].set_title("Drawdown measures loss relative to the previous running peak", loc="left")

    fig.suptitle("Bitcoin: price, volatility and drawdown over 15 years", color=NAVY, weight="bold", fontsize=21)
    fig.text(
        0.01,
        0.005,
        "Source: Coin Metrics Community API, PriceUSD, daily. Volatility uses log returns and sqrt(365).",
        color=GRAY,
        fontsize=9,
    )
    fig.savefig(figure_dir / "day08_btc_15y_price_volatility.png", dpi=220, bbox_inches="tight")
    fig.savefig(figure_dir / "day08_btc_15y_price_volatility.pdf", bbox_inches="tight")
    plt.close(fig)


def plot_annual_volatility(annual: pd.DataFrame, figure_dir: Path) -> None:
    configure_style()
    plot_data = annual.loc[annual["complete_year"]].copy()
    fig, ax = plt.subplots(figsize=(12, 6), constrained_layout=True)
    colours = [
        CRIMSON if value >= plot_data["annualised_volatility"].median() else TEAL
        for value in plot_data["annualised_volatility"]
    ]
    sns.barplot(
        data=plot_data,
        x="year",
        y="annualised_volatility",
        hue="year",
        palette=colours,
        legend=False,
        ax=ax,
    )
    ax.yaxis.set_major_formatter(lambda value, _: f"{100 * value:.0f}%")
    ax.set(xlabel="Year", ylabel="Annualised volatility", title="Bitcoin volatility varies substantially across calendar years")
    ax.tick_params(axis="x", rotation=45)
    ax.text(
        0,
        -0.21,
        "Annualised standard deviation of daily log returns; sqrt(365) scaling.",
        transform=ax.transAxes,
        color=GRAY,
        fontsize=10,
    )
    fig.savefig(figure_dir / "day08_btc_annual_volatility.png", dpi=220, bbox_inches="tight")
    fig.savefig(figure_dir / "day08_btc_annual_volatility.pdf", bbox_inches="tight")
    plt.close(fig)


def main() -> None:
    args = parse_args()
    if args.refresh:
        download_price_history(args.input)
    if not args.input.exists():
        raise FileNotFoundError(f"Missing {args.input}. Run with --refresh or supply --input.")

    args.output_dir.mkdir(parents=True, exist_ok=True)
    prices = load_prices(args.input)
    daily = calculate_metrics(prices)
    yearly = annual_summary(daily)
    summary = headline_summary(daily)

    daily.to_csv(args.output_dir / "day08_btc_15y_daily.csv", index=False)
    yearly.to_csv(args.output_dir / "day08_btc_15y_annual_summary.csv", index=False)
    summary.to_csv(args.output_dir / "day08_btc_15y_summary.csv", index=False)
    plot_history(daily, args.figure_dir)
    plot_annual_volatility(yearly, args.figure_dir)

    print(summary.to_string(index=False))
    print("\nVolatility comparison:")
    print(compare_volatility_windows(daily).map(lambda value: f"{value:.1%}").to_string())
    print(f"\nWrote analysis data to {args.output_dir}")
    print(f"Wrote figures to {args.figure_dir}")


if __name__ == "__main__":
    main()
