"""Generate two small, wholly synthetic teaching examples. Python stdlib only."""

from __future__ import annotations

import hashlib
import json
import math
from pathlib import Path


ROOT = Path(__file__).resolve().parent
RETURNS = (4, 2, -2, 0)  # percent; chosen by hand for the published four-period example
GROSS = 0.40  # percent of a common notional
COMMISSION = 0.05  # percent per fill
SPREAD = 0.02  # percent allocation per fill
SLIPPAGE = 0.03  # percent per fill


def metrics() -> dict[str, float]:
    mean = sum(RETURNS) / len(RETURNS)
    population_sd = math.sqrt(sum((value - mean) ** 2 for value in RETURNS) / len(RETURNS))
    downside = math.sqrt(sum(min(0, value) ** 2 for value in RETURNS) / len(RETURNS))
    total_cost = 2 * (COMMISSION + SPREAD + SLIPPAGE)
    return {
        "mean_pct": mean,
        "population_sd_pct": population_sd,
        "downside_deviation_pct": downside,
        "sharpe_unannualized": mean / population_sd,
        "sortino_unannualized": mean / downside,
        "round_trip_cost_pct": total_cost,
        "net_trade_pct": GROSS - total_cost,
    }


def returns_csv() -> str:
    return "period,return_pct\n" + "".join(f"{i},{value}\n" for i, value in enumerate(RETURNS, 1))


def costs_csv() -> str:
    return ("leg,gross_pct,commission_pct,spread_allocation_pct,slippage_pct\n"
            f"entry,{GROSS:.2f},{COMMISSION:.2f},{SPREAD:.2f},{SLIPPAGE:.2f}\n"
            f"exit,0.00,{COMMISSION:.2f},{SPREAD:.2f},{SLIPPAGE:.2f}\n")


def returns_svg() -> str:
    bars = "".join(
        f'<rect x="{95 + 100 * i}" y="{180 - max(value, 0) * 25}" width="52" '
        f'height="{abs(value) * 25}" fill="{["#444444", "#777777", "#aaaaaa", "#dddddd"][i]}"/>'
        f'<text x="{121 + 100 * i}" y="{250 if value < 0 else 165 - value * 25}" text-anchor="middle">{value}%</text>'
        f'<text x="{121 + 100 * i}" y="290" text-anchor="middle">{i + 1}</text>'
        for i, value in enumerate(RETURNS)
    )
    return (f'<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 520 330" role="img" '
            f'aria-label="Four synthetic period returns: 4, 2, minus 2 and 0 percent">'
            f'<rect width="520" height="330" fill="#ffffff"/>'
            f'<g fill="#222222" font-family="IBM Plex Sans, sans-serif" font-size="16">'
            f'<text x="25" y="30">Synthetic period returns (%)</text>'
            f'<path d="M70 180H485" stroke="#555555" fill="none"/>{bars}'
            f'<text x="25" y="315">Period</text></g></svg>\n')


def costs_svg() -> str:
    # Widths encode 0.40 gross, 0.20 modeled costs and 0.20 net, not market performance.
    return ('<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 640 230" role="img" '
            'aria-label="Synthetic round trip: gross 0.40 percent, modeled cost 0.20 percent, net 0.20 percent">'
            '<rect width="640" height="230" fill="#ffffff"/>'
            '<g fill="#222222" font-family="IBM Plex Sans, sans-serif" font-size="16">'
            '<text x="20" y="32">One synthetic round trip, percent of notional</text>'
            '<text x="20" y="85">Gross</text><rect x="155" y="65" width="400" height="28" fill="#555555"/>'
            '<text x="565" y="85">0.40%</text>'
            '<text x="20" y="138">Cost</text><rect x="155" y="118" width="200" height="28" fill="#aaaaaa"/>'
            '<text x="365" y="138">0.20%</text>'
            '<text x="20" y="191">Net</text><rect x="155" y="171" width="200" height="28" fill="#333333"/>'
            '<text x="365" y="191">0.20%</text></g></svg>\n')


def generated_files() -> dict[str, bytes]:
    values = metrics()
    payload = {
        "provenance": "Hand-chosen synthetic teaching values; no market, customer or provider data",
        "license": "CC0-1.0",
        "returns_percent": list(RETURNS),
        "assumptions": {"risk_free_pct": 0, "mar_pct": 0, "population_sd": True, "annualized": False,
                        "fills": 2, "common_notional": True},
        "metrics": {key: round(value, 9) for key, value in values.items()},
    }
    return {
        "returns.csv": returns_csv().encode("utf-8"),
        "costs.csv": costs_csv().encode("utf-8"),
        "returns.svg": returns_svg().encode("utf-8"),
        "costs.svg": costs_svg().encode("utf-8"),
        "results.json": (json.dumps(payload, indent=2, sort_keys=True) + "\n").encode("utf-8"),
    }


def write(root: Path = ROOT) -> dict[str, str]:
    files = generated_files()
    for name, content in files.items():
        (root / name).write_bytes(content)
    hashes = {name: hashlib.sha256(content).hexdigest() for name, content in files.items()}
    (root / "SHA256SUMS.json").write_text(json.dumps(hashes, indent=2, sort_keys=True) + "\n", encoding="utf-8")
    return hashes


if __name__ == "__main__":
    write()
