# PythonLand 投資×Python シリーズ 第14回
# 「定期実行で株価を毎日ためる（タスクスケジューラ／cron × SQLite）」
# https://pythonland.tech/stock-data-task-scheduler.html
#
# 1回起動すると「直近の日足を取りにいく → SQLite に UPSERT する → ログを1行残す →
# 終了コードを返して終わる」だけのスクリプトです（2026-09-10 時点）。
# 常駐しません。「毎日」の担当は OS 側（Windows のタスクスケジューラ／Linux の cron）です。
#
# ⚠️ これは投資判断をするプログラムではありません。決まった時刻に値を取ってきて、
#    自分のPCの中のファイルに積むだけの道具です。ためた値が高いか安いか、
#    上がるか下がるかは、このスクリプトは何も判断しませんし、判断材料も示しません。
# ⚠️ 通信について: --demo を付けたときは完全にオフラインで動きます（架空の値を生成するだけ）。
#    それ以外のときだけ yfinance 経由で Yahoo! Finance と通信します。送るのは銘柄コードだけで、
#    DBの中身は送信しません。yfinance は Yahoo! Finance の非公式ラッパーで、
#    公式リポジトリには「研究・教育目的」「Yahoo! Finance の API は個人利用を想定」とあります。
#    取得したデータの再配布・商用利用は Yahoo! の利用規約をご自身で確認してください。
# ⚠️ --demo が生成する SMPL-01 〜 SMPL-03 の値はすべて架空です。実在の銘柄とは無関係で、
#    推奨銘柄でも運用実績でもありません。--codes の例に置いた ^N225 / 1306.T は
#    動作確認のための例示であり、推奨銘柄ではありません。
# ⚠️ 作られる prices.db は「自分がどの銘柄を見ているか」の記録です。共有フォルダ・
#    クラウド同期フォルダ・Webサーバーの公開ディレクトリには置かないでください。
# ライセンス: MIT（https://pythonland.tech/ のコードは MIT ライセンスで公開しています）

"""日足を1回だけ取得して SQLite に UPSERT する定期実行用スクリプト。

使い方:
    python stock_daily_fetch.py --demo                       # 通信せず架空データで動かす
    python stock_daily_fetch.py --codes ^N225,1306.T         # ここだけ外部と通信する
    python stock_daily_fetch.py --codes 1306.T --period 1mo  # 取り直す期間を変える
    python stock_daily_fetch.py --show 8                     # DBの件数と直近の行を表示する
    python stock_daily_fetch.py --chart chart.png            # ためた値を折れ線PNGにする
    python stock_daily_fetch.py --print-schtasks             # Windows の登録コマンド例を表示
    python stock_daily_fetch.py --print-cron                 # Linux の crontab 行の例を表示

終了コード: 0 = 1件以上保存できた／1 = 1件も保存できなかった（取得失敗・空）。
--print-* と --show は表示するだけで、取得も保存も登録もしません。
"""
from __future__ import annotations

import argparse
import logging
import random
import sqlite3
import sys
import time
from datetime import date, timedelta
from pathlib import Path

# 実行時のカレントディレクトリに依存しないよう、置き場所を基準に絶対パス化する。
# タスクスケジューラも cron も、スクリプトのある場所をカレントにしてくれない。
BASE_DIR = Path(__file__).resolve().parent
DEFAULT_DB = BASE_DIR / "prices.db"
DEFAULT_LOG = BASE_DIR / "logs" / "fetch.log"

DEMO_CODES = ("SMPL-01", "SMPL-02", "SMPL-03")
DEMO_ANCHOR = date(2026, 1, 5)   # 架空データの起点（固定＝何度実行しても同じ値になる）

CREATE_SQL = """
CREATE TABLE IF NOT EXISTS prices (
    code       TEXT    NOT NULL,
    date       TEXT    NOT NULL,
    open       REAL,
    high       REAL,
    low        REAL,
    close      REAL    NOT NULL,
    volume     INTEGER,
    fetched_at TEXT    NOT NULL,
    PRIMARY KEY (code, date)
)
"""

UPSERT_SQL = """
INSERT INTO prices (code, date, open, high, low, close, volume, fetched_at)
VALUES (?, ?, ?, ?, ?, ?, ?, datetime('now','localtime'))
ON CONFLICT(code, date) DO UPDATE SET
    open       = excluded.open,
    high       = excluded.high,
    low        = excluded.low,
    close      = excluded.close,
    volume     = excluded.volume,
    fetched_at = excluded.fetched_at
"""

log = logging.getLogger("stock_daily_fetch")


# ---------------------------------------------------------------- ログ
def setup_logging(log_path: Path, quiet: bool = False) -> None:
    """ファイルと画面の両方に出す。スケジューラ実行では画面側は誰も見ない。"""
    log_path.parent.mkdir(parents=True, exist_ok=True)
    log.setLevel(logging.INFO)
    log.handlers.clear()
    fmt = logging.Formatter("%(asctime)s [%(levelname)s] %(message)s",
                            datefmt="%Y-%m-%d %H:%M:%S")
    fh = logging.FileHandler(log_path, encoding="utf-8")
    fh.setFormatter(fmt)
    log.addHandler(fh)
    if not quiet:
        sh = logging.StreamHandler(sys.stdout)
        sh.setFormatter(fmt)
        log.addHandler(sh)


# ---------------------------------------------------------------- DB
def open_db(db_path: Path) -> sqlite3.Connection:
    db_path.parent.mkdir(parents=True, exist_ok=True)
    conn = sqlite3.connect(str(db_path))
    conn.execute(CREATE_SQL)
    conn.commit()
    return conn


def count_rows(conn: sqlite3.Connection) -> int:
    return conn.execute("SELECT COUNT(*) FROM prices").fetchone()[0]


def save_rows(conn: sqlite3.Connection, rows: list[tuple]) -> tuple[int, int]:
    """UPSERT で流し込み、(新規n件, 更新m件) を返す。

    executemany の rowcount では新規と更新を区別できないので、
    保存の前後で行数を数えて差し引く。
    """
    if not rows:
        return 0, 0
    before = count_rows(conn)
    with conn:                       # 例外なく抜けたらコミット、落ちたらロールバック
        conn.executemany(UPSERT_SQL, rows)
    after = count_rows(conn)
    inserted = after - before
    return inserted, len(rows) - inserted


# ---------------------------------------------------------------- 取得
def fetch_history(code: str, period: str, retries: int = 2, timeout: int = 30):
    """1銘柄ぶんの日足を返す。取得できなければ None。

    第1回 fetch_latest_price() と同じ形（空データと例外を分ける／例外はログに残す／
    失敗は None に揃える）。返すものを「最新の終値」から「DataFrame」に変えただけ。
    """
    import yfinance as yf

    for attempt in range(retries + 1):
        try:
            df = yf.Ticker(code).history(period=period, timeout=timeout)
            if df is not None and not df.empty:
                return df
            log.warning("%s: 空のデータが返りました（銘柄コードを確認）", code)
            return None                      # 空はリトライしても変わらない
        except Exception as e:               # noqa: BLE001 1銘柄の失敗で全体を止めない
            log.warning("%s: %s %s", code, type(e).__name__, e)
            if attempt < retries:
                time.sleep(2 ** attempt)     # 1秒 → 2秒と待ち時間を伸ばす
    return None


def fetch_batch(codes: list[str], period: str, timeout: int = 30):
    """まず一括取得を試す。1リクエストで全銘柄ぶんが返る。

    戻り値は {銘柄コード: DataFrame}。取れなかった銘柄はキーごと入らない。
    """
    import pandas as pd
    import yfinance as yf

    out: dict[str, "pd.DataFrame"] = {}
    try:
        raw = yf.download(codes, period=period, group_by="ticker",
                          progress=False, threads=False, timeout=timeout)
    except Exception as e:                   # noqa: BLE001
        log.warning("一括取得に失敗しました: %s %s", type(e).__name__, e)
        return out
    if raw is None or raw.empty:
        log.warning("一括取得の結果が空でした")
        return out

    for code in codes:
        try:
            if isinstance(raw.columns, pd.MultiIndex):
                if code not in raw.columns.get_level_values(0):
                    continue
                df = raw[code]
            else:
                df = raw                     # 1銘柄だけのときは単層で返る版もある
            df = df.dropna(how="all")
            if not df.empty:
                out[code] = df
        except Exception as e:               # noqa: BLE001
            log.warning("%s: 取り出しに失敗 %s %s", code, type(e).__name__, e)
    return out


def frame_to_rows(code: str, df) -> list[tuple]:
    """DataFrame を (code, date, open, high, low, close, volume) の並びに変換する。

    日付は「実行した日」ではなく DataFrame の index（取引日）を使う。
    """
    import pandas as pd

    rows = []
    for idx, row in df.iterrows():
        day = idx.date().isoformat() if hasattr(idx, "date") else str(idx)[:10]
        close = row.get("Close")
        if close is None or pd.isna(close):
            continue                          # 終値が無い行は積まない

        def num(name):
            v = row.get(name)
            return None if v is None or pd.isna(v) else float(v)

        vol = row.get("Volume")
        rows.append((code, day, num("Open"), num("High"), num("Low"),
                     float(close), None if vol is None or pd.isna(vol) else int(vol)))
    return rows


# ---------------------------------------------------------------- 架空データ
def demo_rows(codes: list[str], days: int = 30) -> list[tuple]:
    """完全オフラインで架空の日足を作る。何度実行しても同じ値になる。

    ⚠️ ここで作られる値は架空です。実在の銘柄・実在の値動きとは無関係です。
    """
    rows = []
    today = date.today()
    for i, code in enumerate(codes):
        price = 1000.0 + 500.0 * i
        day = DEMO_ANCHOR
        series = []
        while day <= today:
            if day.weekday() < 5:             # 平日だけ（架空の「取引日」）
                rnd = random.Random(f"{code}:{day.isoformat()}")
                price *= 1.0 + rnd.uniform(-0.02, 0.02)
                high = price * (1.0 + rnd.uniform(0.0, 0.01))
                low = price * (1.0 - rnd.uniform(0.0, 0.01))
                opn = low + (high - low) * rnd.random()
                series.append((code, day.isoformat(), round(opn, 1), round(high, 1),
                               round(low, 1), round(price, 1), rnd.randrange(1000, 9000) * 100))
            day += timedelta(days=1)
        rows.extend(series[-days:])
    return rows


# ---------------------------------------------------------------- 表示系
def show_db(conn: sqlite3.Connection, limit: int) -> None:
    total = count_rows(conn)
    print(f"件数: {total}")
    cur = conn.execute(
        "SELECT code, date, close, fetched_at FROM prices "
        "ORDER BY date DESC, code ASC LIMIT ?", (limit,))
    print(f"{'code':<10}{'date':<12}{'close':>10}  fetched_at")
    for code, day, close, fetched in cur:
        print(f"{code:<10}{day:<12}{close:>10.1f}  {fetched}")


def draw_chart(conn: sqlite3.Connection, out_path: Path) -> None:
    """ためた終値を折れ線PNGにする。画面のない環境でも同じ絵が出る。"""
    import matplotlib
    matplotlib.use("Agg")            # 画面を開かない描画方式。import の前に指定する
    import matplotlib.pyplot as plt

    codes = [r[0] for r in conn.execute("SELECT DISTINCT code FROM prices ORDER BY code")]
    if not codes:
        print("DBが空です。先に取得してください。")
        return
    fig, ax = plt.subplots(figsize=(9, 4.5))
    for code in codes:
        cur = conn.execute("SELECT date, close FROM prices WHERE code = ? ORDER BY date",
                           (code,))
        pairs = cur.fetchall()
        ax.plot([p[0] for p in pairs], [p[1] for p in pairs], label=code, linewidth=1.6)
    ax.set_title("saved close prices in prices.db")
    ax.set_xlabel("trading date")
    ax.set_ylabel("close")
    ax.legend()
    ax.grid(alpha=0.3)
    step = max(1, len(ax.get_xticks()) // 10)
    for i, label in enumerate(ax.get_xticklabels()):
        label.set_visible(i % step == 0)
        label.set_rotation(45)
        label.set_horizontalalignment("right")
    fig.tight_layout()
    out_path.parent.mkdir(parents=True, exist_ok=True)
    fig.savefig(out_path, dpi=110)
    plt.close(fig)
    print(f"保存しました: {out_path}")


def print_schtasks(args) -> None:
    """Windows の登録コマンド例を組み立てて表示するだけ（登録はしない）。"""
    py = Path(sys.executable)
    script = Path(__file__).resolve()
    bat = script.parent / "run_fetch.bat"
    codes = args.codes or "^N225,1306.T"
    print("=== 1. まず .bat を作る（作業ディレクトリと文字コードをここで固定する）===")
    print(f"ファイル: {bat}")
    print("@echo off")
    print('cd /d "%~dp0"')
    # cmd.exe では ^ がエスケープ文字なので、指数の ^N225 は必ず引用符で囲む
    print(f'"{py}" -X utf8 "{script.name}" --codes "{codes}"')
    print("exit /b %ERRORLEVEL%")
    print()
    print("=== 2. タスクを登録する（cmd.exe で1行。^ は継続行）===")
    print(f'schtasks /create /tn "\\PythonLand\\StockDaily" /tr "\'{bat}\'" '
          f"/sc DAILY /st 18:30 /f")
    print()
    print("=== 3. 確認・テスト・削除 ===")
    print('schtasks /query /tn "\\PythonLand\\StockDaily" /v /fo LIST')
    print('schtasks /run   /tn "\\PythonLand\\StockDaily"')
    print('schtasks /delete /tn "\\PythonLand\\StockDaily" /f')
    print()
    print("※ 表示するだけです。実際の登録・削除はご自身で確認してから実行してください。")


def print_cron(args) -> None:
    """Linux の crontab 行の例を組み立てて表示するだけ（登録はしない）。"""
    py = Path(sys.executable)
    script = Path(__file__).resolve()
    logfile = Path(args.log).resolve() if args.log else DEFAULT_LOG
    codes = args.codes or "^N225,1306.T"
    print("=== crontab -e に貼る1行（平日18:30・すべて絶対パス）===")
    print(f"30 18 * * 1-5 {py} {script} --codes '{codes}' "
          f">> {logfile.parent / 'cron.log'} 2>&1")
    print()
    print("=== 確認・削除 ===")
    print("crontab -l            # 今の内容を表示（-r は全消しなので押し間違いに注意）")
    print()
    print("※ 表示するだけです。cron は指定時刻にPCが動いていることが前提です。")


# ---------------------------------------------------------------- 本体
def build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(
        description="日足を1回だけ取得して SQLite に UPSERT する（常駐しません）",
        formatter_class=argparse.RawDescriptionHelpFormatter)
    p.add_argument("--codes", help="銘柄コードをカンマ区切りで（例: ^N225,1306.T）")
    p.add_argument("--period", default="5d",
                   help="1回の実行で取り直す期間（既定: 5d＝直近5営業日）")
    p.add_argument("--db", default=str(DEFAULT_DB), help=f"DBファイル（既定: {DEFAULT_DB}）")
    p.add_argument("--log", default=str(DEFAULT_LOG), help=f"ログファイル（既定: {DEFAULT_LOG}）")
    p.add_argument("--sleep", type=float, default=1.0,
                   help="銘柄ごとに取りにいくときの待ち時間（秒・既定 1.0）")
    p.add_argument("--demo", action="store_true",
                   help="通信せず架空データで動かす（SMPL-01〜03）")
    p.add_argument("--quiet", action="store_true", help="画面には出さずログだけに書く")
    p.add_argument("--show", type=int, metavar="N", help="DBの件数と直近N行を表示して終わる")
    p.add_argument("--chart", metavar="PNG", help="ためた終値を折れ線PNGにして終わる")
    p.add_argument("--print-schtasks", action="store_true",
                   help="Windows の登録コマンド例を表示して終わる")
    p.add_argument("--print-cron", action="store_true",
                   help="Linux の crontab 行の例を表示して終わる")
    return p


def main(argv=None) -> int:
    args = build_parser().parse_args(argv)

    if args.print_schtasks:
        print_schtasks(args)
        return 0
    if args.print_cron:
        print_cron(args)
        return 0

    db_path = Path(args.db).expanduser()
    if args.show is not None or args.chart:
        conn = open_db(db_path)
        try:
            if args.show is not None:
                show_db(conn, args.show)
            if args.chart:
                draw_chart(conn, Path(args.chart).expanduser())
        finally:
            conn.close()
        return 0

    setup_logging(Path(args.log).expanduser(), quiet=args.quiet)
    log.info("開始 db=%s demo=%s period=%s", db_path, args.demo, args.period)

    if args.demo:
        codes = list(DEMO_CODES)
        rows = demo_rows(codes)
        failed: list[str] = []
    else:
        if not args.codes:
            log.error("--codes か --demo のどちらかを指定してください")
            return 1
        codes = [c.strip() for c in args.codes.split(",") if c.strip()]
        log.info("対象 %d件: %s", len(codes), ", ".join(codes))
        frames = fetch_batch(codes, args.period)
        failed = [c for c in codes if c not in frames]
        for code in list(failed):                 # 一括で取れなかったぶんだけ個別に
            time.sleep(args.sleep)
            df = fetch_history(code, args.period)
            if df is not None and not df.empty:
                frames[code] = df
                failed.remove(code)
        rows = []
        for code in codes:
            if code in frames:
                rows.extend(frame_to_rows(code, frames[code]))

    if not rows:
        log.error("終了 保存0件 exit=1（取得できた行がありません）")
        return 1

    conn = open_db(db_path)
    try:
        inserted, updated = save_rows(conn, rows)
        total = count_rows(conn)
    finally:
        conn.close()

    if failed:
        log.warning("取得できなかった銘柄 %d件: %s", len(failed), ", ".join(failed))
    log.info("保存 %d件（新規 %d / 更新 %d）DB合計 %d件", len(rows), inserted, updated, total)
    log.info("終了 exit=0")
    return 0


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