# PythonLand 投資×Python シリーズ 第15回
# 「ポートフォリオのリスク分析（標準偏差・相関）」
# https://pythonland.tech/portfolio-risk-analysis.html
#
# 貯めておいた終値を読んで、日次リターンの「散らばり（標準偏差）」と
# 「銘柄同士の連動（相関係数）」を数値にして表示するだけのスクリプトです（2026-09-16 時点）。
#
# ⚠️ これは投資判断をするプログラムではありません。出すのは過去の値動きの散らばりを表す
#    数値だけで、その銘柄が良いか悪いか、安全か危険か、これからどうなるかは一切扱いません。
#    配分（ウェイト）の最適化もしません。--weights は「あなたが決めた配分」を受け取るだけで、
#    既定値を持ちません（既定値を置くと「おすすめの配分」に見えてしまうため）。
# ⚠️ 通信は1回も発生しません。読むのは手元の SQLite（読み取り専用で開きます）か CSV か、
#    --demo で自分が生成した架空データだけです。株価を取りに行く処理は入っていません。
# ⚠️ DB は file:...?mode=ro で開くので、このスクリプトは prices.db に1行も書き込みません。
#    第8回の dividends.db・第11回の perks.db には接続もしません。
# ⚠️ --demo が生成する SMPL-A 〜 SMPL-E の値はすべて架空です。実在の銘柄とは無関係で、
#    推奨銘柄でも運用実績でもありません。
# ライセンス: MIT（https://pythonland.tech/ のコードは MIT ライセンスで公開しています）

"""貯めた終値から標準偏差・相関行列・合成ポートフォリオの標準偏差を出す CLI。

使い方:
    python portfolio_risk.py --demo                          # 架空データ250営業日で動かす
    python portfolio_risk.py --demo --ragged                 # 日付が揃っていない架空データ
    python portfolio_risk.py --db prices.db                  # 第14回で貯めたDBを読む（読み取り専用）
    python portfolio_risk.py --csv prices.csv                # date,code,close の3列CSVを読む
    python portfolio_risk.py --demo --weights 0.5,0.3,0.2 --codes SMPL-A,SMPL-B,SMPL-C
    python portfolio_risk.py --demo --heatmap-compare heat.png   # 相関行列の図（2枚並べ）
    python portfolio_risk.py --demo --spread spread.png --weights 0.5,0.5 --codes SMPL-A,SMPL-C
    python portfolio_risk.py --demo-split                    # 調整基準が混ざったDBの再現実験
    python portfolio_risk.py --demo --write-csv prices.csv   # 架空データをCSVに書き出す

終了コード: 0 = 表を出せた／1 = 入力が足りない（列が2本未満・行が足りない等）。
"""
from __future__ import annotations

import argparse
import sqlite3
import sys
from pathlib import Path

import numpy as np
import pandas as pd

import matplotlib
matplotlib.use("Agg")          # 画面を開かない。PNGに書き出すだけ（サーバーでも動く）
import matplotlib.pyplot as plt  # noqa: E402  （use() の後に import するのが作法）

BASE_DIR = Path(__file__).resolve().parent
DEMO_CODES = ("SMPL-A", "SMPL-B", "SMPL-C", "SMPL-D")
DEMO_START = "2025-01-06"      # 架空データの起点（固定＝何度実行しても同じ値になる）


# ---------------------------------------------------------------- 入力（3通り）

def load_from_db(db_path: Path, table: str = "prices") -> pd.DataFrame:
    """SQLite を読み取り専用で開き、date × code の横持ち表にして返す。

    mode=ro は「書き込もうとした時点でエラーになる」モード。うっかり INSERT を
    書いてしまっても、他人のDBを壊さずに済む（第14回で作った prices.db を守るため）。
    """
    uri = db_path.resolve().as_uri() + "?mode=ro"
    with sqlite3.connect(uri, uri=True) as conn:
        df = pd.read_sql_query(
            f"SELECT code, date, close FROM {table} ORDER BY date, code", conn)
    return _to_wide(df)


def load_from_csv(csv_path: Path) -> pd.DataFrame:
    """date,code,close の3列CSVを読む（証券会社や表計算から出したものを想定）。"""
    df = pd.read_csv(csv_path)
    missing = {"date", "code", "close"} - set(df.columns)
    if missing:
        raise SystemExit(f"CSVに列がありません: {sorted(missing)}（必要: date,code,close）")
    return _to_wide(df)


def _to_wide(df: pd.DataFrame) -> pd.DataFrame:
    """縦持ち（1行=1銘柄1日）を横持ち（index=日付・列=銘柄）にする。"""
    df = df.copy()
    df["date"] = pd.to_datetime(df["date"])
    df["close"] = pd.to_numeric(df["close"], errors="coerce")
    wide = df.pivot(index="date", columns="code", values="close").sort_index()
    wide.columns.name = None
    return wide


def make_demo(days: int = 250, seed: int = 42, ragged: bool = False) -> pd.DataFrame:
    """架空の終値表を作る（通信なし）。

    乱数はシード固定なので、何度実行しても同じ値が出ます。ただし NumPy の
    バージョンが変わると生成される乱数列そのものが変わることがあります
    （同じ版なら同じ、が保証できる範囲です）。
    """
    rng = np.random.default_rng(seed)
    dates = pd.bdate_range(DEMO_START, periods=days)
    target_corr = np.array([[1.0, 0.9, 0.1, -0.3],
                            [0.9, 1.0, 0.1, -0.2],
                            [0.1, 0.1, 1.0, 0.0],
                            [-0.3, -0.2, 0.0, 1.0]])
    vols = np.array([0.015, 0.012, 0.020, 0.018])       # 1日あたりの散らばりの目標値
    logret = rng.multivariate_normal(np.zeros(4), np.outer(vols, vols) * target_corr,
                                     size=days)
    px = pd.DataFrame(1000 * np.exp(np.cumsum(logret, axis=0)),
                      index=dates, columns=list(DEMO_CODES)).round(1)
    if not ragged:
        return px
    # ここから下は「貯めたDBあるある」を再現する穴あけ（架空）
    px = px.copy()
    px.iloc[:days - 100, px.columns.get_loc("SMPL-B")] = np.nan   # 途中から取り始めた銘柄
    px.iloc[::7, px.columns.get_loc("SMPL-C")] = np.nan           # 取りこぼした日がある銘柄
    px["SMPL-E"] = np.nan
    px.iloc[10:13, px.columns.get_loc("SMPL-E")] = [1000.0, 1010.0, 1005.0]  # 3日だけの銘柄
    return px


# ------------------------------------------------------------ 計算（統計量だけ）

def daily_returns(px: pd.DataFrame, log: bool = False) -> pd.DataFrame:
    """日次リターン（前日比の変化率）。欠損は埋めずに NaN のまま残す。

    pandas 3.0 では pct_change() は前埋めをしません（fill_method は None 必須）。
    明示的に書いているのは「埋めていない」と読む人に伝えるためです。
    """
    if log:
        return np.log(px).diff()
    return px.pct_change(fill_method=None)


def stats_table(ret: pd.DataFrame, days_per_year: int) -> pd.DataFrame:
    """銘柄ごとに、本数・標準偏差（ddof=1 と 0）・年率換算を並べた表を作る。"""
    n = ret.notna().sum()
    sd1 = ret.std()                 # pandas の既定は ddof=1（N-1で割る）
    sd0 = ret.std(ddof=0)           # numpy の既定と同じ（Nで割る）
    k = np.sqrt(days_per_year)
    return pd.DataFrame({
        "本数": n,
        "std(ddof=1)": sd1.round(8),
        "std(ddof=0)": sd0.round(8),
        f"年率%(ddof=1)": (sd1 * k * 100).round(4),
        f"年率%(ddof=0)": (sd0 * k * 100).round(4),
    })


def overlap_matrix(ret: pd.DataFrame) -> pd.DataFrame:
    """ペアごとに「両方とも値がある日」が何日あったかを数える。

    相関係数は必ずこの表とセットで見ます。0.98 でも重なり2日なら意味が違うからです。
    """
    ok = ret.notna().astype(int)
    return ok.T @ ok


def composite_sd(ret: pd.DataFrame, weights: np.ndarray, ddof: int = 1) -> dict:
    """合成ポートフォリオの標準偏差を、行列の式と系列の2通りで出す（検算つき）。

    ⚠️ ここで返すのは「その配分だったら、過去の散らばりはいくつだったか」だけです。
       良い配分・最適な配分を探す処理は入れていません。
    """
    used = ret.dropna()                              # 全銘柄そろった日だけで計算する
    cov = used.cov(ddof=ddof)
    sd_matrix = float(np.sqrt(weights @ cov.to_numpy() @ weights))
    sd_series = float((used.to_numpy() @ weights).std(ddof=ddof))
    sd_each = used.std(ddof=ddof).to_numpy()
    return {
        "rows": len(used),
        "sd_matrix": sd_matrix,
        "sd_series": sd_series,
        "diff": abs(sd_matrix - sd_series),
        "sd_weighted_avg": float(sd_each @ weights),
        "ratio": sd_matrix / float(sd_each @ weights),
    }


def big_moves(ret: pd.DataFrame, threshold: float) -> pd.DataFrame:
    """1日で threshold を超えて動いた行を拾う（調整基準の混在を疑うための警告）。"""
    hits = ret.abs() > threshold
    if not hits.to_numpy().any():
        return pd.DataFrame(columns=["date", "code", "return"])
    idx, col = np.where(hits.to_numpy())
    return pd.DataFrame({
        "date": ret.index[idx].strftime("%Y-%m-%d"),
        "code": ret.columns[col],
        "return": ret.to_numpy()[idx, col].round(6),
    })


# ---------------------------------------------------------------------- 図（2枚）

def draw_heatmap(ax, corr: pd.DataFrame, fix_scale: bool, title: str) -> None:
    """相関行列を imshow だけで描く（seaborn を使わない）。"""
    kw = {"vmin": -1, "vmax": 1} if fix_scale else {}
    im = ax.imshow(corr.to_numpy(), cmap="coolwarm", **kw)
    ax.set_xticks(range(len(corr.columns)), corr.columns, rotation=45, ha="right")
    ax.set_yticks(range(len(corr.index)), corr.index)
    ax.set_title(title, fontsize=10)
    for i in range(corr.shape[0]):
        for j in range(corr.shape[1]):
            v = corr.to_numpy()[i, j]
            ax.text(j, i, "nan" if np.isnan(v) else f"{v:.2f}",
                    ha="center", va="center", fontsize=9, color="#222")
    cb = ax.figure.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
    if fix_scale:
        cb.set_ticks([-1, -0.5, 0, 0.5, 1])


def save_heatmap(corr: pd.DataFrame, path: Path, compare: bool = False) -> None:
    # 図の中の文字は英字だけにしている（matplotlib 同梱フォントに日本語が無く□になる）
    if compare:
        fig, axes = plt.subplots(1, 2, figsize=(11, 4.6), dpi=110)
        draw_heatmap(axes[0], corr, False, "no vmin/vmax (default)")
        draw_heatmap(axes[1], corr, True, "vmin=-1, vmax=1 (fixed)")
    else:
        fig, ax = plt.subplots(figsize=(5.6, 4.6), dpi=110)
        draw_heatmap(ax, corr, True, "correlation (vmin=-1, vmax=1)")
    fig.suptitle("daily return correlation / demo data (fictional)", fontsize=10)
    fig.tight_layout()
    fig.savefig(path)
    plt.close(fig)


def save_spread(ret: pd.DataFrame, weights: np.ndarray, path: Path, ddof: int = 1) -> None:
    """個別と合成の日次リターンの散らばりをヒストグラムで重ねる。"""
    used = ret.dropna()
    fig, ax = plt.subplots(figsize=(7.4, 4.4), dpi=110)
    bins = np.linspace(-0.06, 0.06, 61)
    for code in used.columns:
        ax.hist(used[code], bins=bins, histtype="step", linewidth=1.2,
                label=f"{code}  sd={used[code].std(ddof=ddof):.5f}")
    mix = used.to_numpy() @ weights
    ax.hist(mix, bins=bins, histtype="stepfilled", alpha=0.35, color="#444",
            label=f"mix {np.round(weights, 3).tolist()}  sd={mix.std(ddof=ddof):.5f}")
    ax.set_xlabel("daily return")
    ax.set_ylabel("days")
    ax.set_title("spread of daily returns / demo data (fictional)", fontsize=10)
    ax.legend(fontsize=8)
    fig.tight_layout()
    fig.savefig(path)
    plt.close(fig)


# ------------------------------------------- 調整基準が混ざったDBの再現実験（架空）

def demo_split(days: int = 40, split_at: int = 20, ratio: float = 2.0,
               dividend: float = 0.0187, seed: int = 7, days_per_year: int = 252) -> None:
    """「分割前に貯めた行だけ古い基準のまま残っているDB」を架空データで再現する。"""
    rng = np.random.default_rng(seed)
    dates = pd.bdate_range("2025-06-02", periods=days)
    true_px = pd.Series((1000 * np.exp(np.cumsum(rng.normal(0, 0.013, days)))).round(1),
                        index=dates)
    db_split = true_px.copy()
    db_split.iloc[:split_at] = (db_split.iloc[:split_at] * ratio).round(1)
    db_div = true_px.copy()
    db_div.iloc[:split_at] = (db_div.iloc[:split_at] / (1 - dividend)).round(1)

    k = np.sqrt(days_per_year)
    r_true = true_px.pct_change(fill_method=None)
    r_split = db_split.pct_change(fill_method=None)
    r_div = db_div.pct_change(fill_method=None)
    d = dates[split_at]
    print(f"[再現実験] 架空の40営業日・{d:%Y-%m-%d} に 1:{ratio:g} の株式分割があったと仮定")
    print(f"  DBに残った値      : {db_split.iloc[split_at - 1]:.1f} -> {db_split.iloc[split_at]:.1f}"
          f"   日次リターン {r_split.iloc[split_at]:+.4f}")
    print(f"  取り直した正しい値: {true_px.iloc[split_at - 1]:.1f} -> {true_px.iloc[split_at]:.1f}"
          f"   日次リターン {r_true.iloc[split_at]:+.4f}")
    print(f"  |日次リターン|>20% の件数  DB {int((r_split.abs() > 0.2).sum())} 件 / "
          f"正しい系列 {int((r_true.abs() > 0.2).sum())} 件")
    print(f"  年率SD  DB {r_split.std() * k * 100:.2f}%  /  正しい系列 {r_true.std() * k * 100:.2f}%")
    print(f"[参考] 分配金 {dividend * 100:.2f}% ぶんの遡及調整だけでも")
    print(f"  境界日の日次リターン {r_div.iloc[split_at]:+.4f}（正しくは {r_true.iloc[split_at]:+.4f}）")
    print(f"  年率SD  DB {r_div.std() * k * 100:.2f}%  /  正しい系列 {r_true.std() * k * 100:.2f}%")
    print("  ※ 値はすべて架空です。実在の銘柄の分割・分配金とは関係ありません。")


# ------------------------------------------------------------------------ CLI

def build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(
        description="貯めた終値から標準偏差・相関行列・合成の標準偏差を表示します"
                    "（投資判断はしません・配分の最適化もしません）")
    src = p.add_mutually_exclusive_group()
    src.add_argument("--db", type=Path, help="SQLiteのパス（読み取り専用で開きます）")
    src.add_argument("--csv", type=Path, help="date,code,close の3列CSV")
    src.add_argument("--demo", action="store_true", help="通信せず架空データを生成する")
    p.add_argument("--table", default="prices", help="DBのテーブル名（既定: prices）")
    p.add_argument("--demo-days", type=int, default=250, help="架空データの営業日数（既定: 250）")
    p.add_argument("--seed", type=int, default=42, help="架空データの乱数シード（既定: 42）")
    p.add_argument("--ragged", action="store_true", help="架空データの日付をわざと揃えない")
    p.add_argument("--codes", help="対象の銘柄コード（カンマ区切り・省略で全部）")
    p.add_argument("--weights", help="合成する配分（カンマ区切り・--codes と同じ順）"
                                     "※ 既定値はありません。自分で決めた値を渡してください")
    p.add_argument("--min-periods", type=int, default=20,
                   help="相関を計算するのに必要な重なり日数（既定: 20）")
    p.add_argument("--days-per-year", type=int, default=252,
                   help="年率換算に使う営業日数の仮定（既定: 252）")
    p.add_argument("--ddof", type=int, default=1, help="標準偏差・共分散の自由度（既定: 1）")
    p.add_argument("--log-return", action="store_true", help="対数リターンで計算する")
    p.add_argument("--dropna", action="store_true", help="全銘柄そろった日だけを使う")
    p.add_argument("--jump", type=float, default=0.2, help="警告する日次リターンの大きさ（既定: 0.2）")
    p.add_argument("--heatmap", type=Path, help="相関行列のPNGを書き出す")
    p.add_argument("--heatmap-compare", type=Path, help="vmin/vmax の有無を並べたPNGを書き出す")
    p.add_argument("--spread", type=Path, help="日次リターンの分布のPNGを書き出す（--weights 必須）")
    p.add_argument("--write-csv", type=Path, help="読み込んだ終値表をCSVに書き出す")
    p.add_argument("--demo-split", action="store_true",
                   help="調整基準が混ざったDBの再現実験だけを表示して終わる")
    return p


def main(argv: list[str] | None = None) -> int:
    args = build_parser().parse_args(argv)
    print(f"[環境] pandas {pd.__version__} / numpy {np.__version__} / "
          f"matplotlib {matplotlib.__version__} / sqlite {sqlite3.sqlite_version}")

    if args.demo_split:
        demo_split(days_per_year=args.days_per_year)
        return 0

    if args.db:
        px = load_from_db(args.db, args.table)
        print(f"[入力] {args.db}（読み取り専用・{args.table} テーブル）")
    elif args.csv:
        px = load_from_csv(args.csv)
        print(f"[入力] {args.csv}")
    elif args.demo:
        px = make_demo(args.demo_days, args.seed, args.ragged)
        print(f"[入力] 架空データ {args.demo_days}営業日 / seed={args.seed}"
              f"{' / 日付を揃えていない版' if args.ragged else ''}")
    else:
        print("入力を指定してください（--db / --csv / --demo のどれか）")
        return 1

    if args.codes:
        wanted = [c.strip() for c in args.codes.split(",") if c.strip()]
        missing = [c for c in wanted if c not in px.columns]
        if missing:
            print(f"入力に無い銘柄コードです: {missing}")
            return 1
        px = px[wanted]

    if args.write_csv:
        long = px.stack().rename("close").reset_index()
        long.columns = ["date", "code", "close"]
        long["date"] = long["date"].dt.strftime("%Y-%m-%d")
        long.to_csv(args.write_csv, index=False)
        print(f"[出力] {args.write_csv} に {len(long)} 行書き出しました")

    print(f"[終値表] {px.shape[0]}日 × {px.shape[1]}銘柄  "
          f"{px.index.min():%Y-%m-%d} 〜 {px.index.max():%Y-%m-%d}")
    if px.shape[1] < 2:
        print("相関を出すには銘柄が2つ以上要ります")
        return 1

    ret = daily_returns(px, log=args.log_return)
    kind = "対数リターン" if args.log_return else "単純リターン"
    if args.dropna:
        before = len(ret.dropna(how="all"))
        ret = ret.dropna()
        print(f"[--dropna] 全銘柄そろった日だけ残しました: {before}日 → {len(ret)}日")

    print(f"\n== 1. 銘柄ごとの散らばり（{kind}・年率は{args.days_per_year}営業日の仮定）==")
    print(stats_table(ret, args.days_per_year).to_string())

    print("\n== 2. ペアごとの重なり日数（この日数で相関が計算されています）==")
    print(overlap_matrix(ret).to_string())

    print(f"\n== 3. 相関行列（pearson・min_periods={args.min_periods}）==")
    corr = ret.corr(min_periods=args.min_periods)
    print(corr.round(4).to_string())
    if corr.isna().to_numpy().any():
        print(f"  ※ NaN のペアは重なりが {args.min_periods} 日未満です（計算していません）")

    print(f"\n== 4. 共分散行列（ddof={args.ddof}）==")
    print(ret.cov(ddof=args.ddof).map(lambda v: f"{v:.3e}").to_string())

    jumps = big_moves(ret, args.jump)
    if len(jumps):
        print(f"\n== 5. 1日で{args.jump * 100:.0f}%以上動いた行（{len(jumps)}件）==")
        print(jumps.to_string(index=False))
        print("  ※ 株式分割や分配金の遡及調整で、古い基準の行が混ざっている可能性があります")

    if args.weights:
        w = np.array([float(x) for x in args.weights.split(",")], dtype=float)
        if len(w) != ret.shape[1]:
            print(f"--weights の個数（{len(w)}）が銘柄数（{ret.shape[1]}）と合いません")
            return 1
        r = composite_sd(ret, w, ddof=args.ddof)
        k = np.sqrt(args.days_per_year)
        print(f"\n== 6. 指定した配分 {np.round(w, 4).tolist()} での合成（{r['rows']}日で計算）==")
        print(f"  行列の式 sqrt(w@cov@w) = {r['sd_matrix']:.10f}")
        print(f"  合成系列の std()        = {r['sd_series']:.10f}   （差 {r['diff']:.2e}）")
        print(f"  各銘柄のsdの加重平均    = {r['sd_weighted_avg']:.10f}")
        print(f"  合成 / 加重平均         = {r['ratio']:.6f}")
        print(f"  年率  合成 {r['sd_matrix'] * k * 100:.3f}%  /  加重平均 "
              f"{r['sd_weighted_avg'] * k * 100:.3f}%")
        print("  ※ 出しているのは過去の散らばりだけです。将来の値動きの予測ではありません。")
        if args.spread:
            save_spread(ret, w, args.spread, ddof=args.ddof)
            print(f"  [出力] {args.spread}")
    elif args.spread:
        print("\n--spread には --weights が必要です（配分の既定値は持っていません）")

    if args.heatmap:
        save_heatmap(corr, args.heatmap)
        print(f"\n[出力] {args.heatmap}")
    if args.heatmap_compare:
        save_heatmap(corr, args.heatmap_compare, compare=True)
        print(f"[出力] {args.heatmap_compare}")
    return 0


if __name__ == "__main__":
    sys.exit(main())
