"""Reconcile the net result of every transaction in the frozen sample.

Accounting basis (per fee-payer wallet):
  value change = native SOL change of the payer
               + for each payer-owned wSOL token account: its lamport change (amount + rent)
               + for each other payer-owned token account: token change (priced if USDC/USDT)
                 plus its lamport change (account rent, recoverable when closed)
Costs are itemised from the transaction itself: base fee, priority fee, Jito tips.
Tokens without a trusted price are reported as "not measured", never as zero.
"""

import collections
import gzip
import json
import os
import statistics
import sys
from decimal import Decimal

EVIDENCE = sys.argv[1] if len(sys.argv) > 1 else "evidence"

WSOL = "So11111111111111111111111111111111111111112"
STABLES = {
    "EPjFWdd5AufqSSqeM2qN1xzybapC8G4wEGGkZwyTDt1v": "USDC",
    "Es9vMFrzaCERmJfrF4H2FYD4KCoNkY11McCe8BenwNYB": "USDT",
}
JITO_TIP_ACCOUNTS = {
    "96gYZGLnJYVFmbjzopPSU6QiEV5fGqZNyN9nmNhvrZU5",
    "HFqU5x63VTqvQss8hp11i4wVV8bD44PvwucfZ2bU7gRe",
    "Cw8CFyM9FkoMi7K7Crf6HNQqf4uEMzpKw6QNghXLvLkY",
    "ADaUMid9yfUytqMBgopwjb2DTLSokTSzL1zt6iGPaS49",
    "DfXygSm4jCyNCybVYYK6DwvWqjKee8pbDmJGcLWNDXjh",
    "ADuUkR4vqLUMWXxW9gh6D6L8pMSawimctcNZ5pGwDcEt",
    "DttWaMuVvTiduZRnguLF7jNxTgiMBZ1hyAumKUiL2KRL",
    "3AVi9Tg9Uo68tJfuvoKvqKNWKkC5wPdSSdeBnizKZ6jT",
}
VENUES = {
    "675kPX9MHTjS2zt1qfr1NYHuzeLXfQM9H24wFSUt1Mp8": "Raydium AMM",
    "CAMMCzo5YL8w4VFF8KVHrK22GGUsp5VTaW7grrKgrWqK": "Raydium CLMM",
    "CPMMoo8L3F4NbTegBCKVNunggL7H1ZpdTHKxQB5qKP1C": "Raydium CPMM",
    "whirLbMiicVdio4qvUfM5KAg6Ct8VwpYzGff3uctyCc": "Orca Whirlpool",
    "LBUZKhRxPF3XUpBCjp4YzTKgLccjZhTSDM9YuVaPwxo": "Meteora DLMM",
    "Eo7WjKq67rjJQSZxS6z3YkapzY3eMj6Xy8X5EQVn5UaB": "Meteora DAMM v1",
    "cpamdpZCGKUy5JxQXB4dcpGPiikHawvSWAd6mEn1sGG": "Meteora DAMM v2",
    "pAMMBay6oceH9fJKBRHGP5D4bD4sWpmSwMn52FMfXEA": "PumpSwap",
    "JUP6LkbZbjS1jKKwapdHNy74zcZ3tLUZoi5QNyVTaV4": "Jupiter v6",
    "PhoeNiXZ8ByJGLkxNfZRnkUfjvmuYqLR89jjFHGqdXY": "Phoenix",
    "SoLFiHG9TfgtdUXUjWAxi3LtvYuFyDLVhBWxdMZxyCe": "SolFi",
    "9H6tua7jkLhdm3w8BvgpTn5LZNU7g4ZynDmCiNN3q6Rp": "HumidiFi",
    "ZERor4xhbUycZ6gb9ntrhqscUcZmAbQDjEAtCf4hbZY": "ZeroFi",
    "TessVdML9pBGgG9yGks7o4HewRaXVAMuoVj4x83GLQH": "Tessera V",
}
BASE_FEE_PER_SIGNATURE = 5000
LAMPORTS = Decimal(10**9)


def token_rows(meta, key):
    return {row["accountIndex"]: row for row in (meta.get(key) or [])}


def amount(row):
    return Decimal(row["uiTokenAmount"]["amount"]) / (Decimal(10) ** row["uiTokenAmount"]["decimals"])


def walk_instructions(tx):
    message = tx["transaction"]["message"]
    yield from message["instructions"]
    for inner in tx["meta"].get("innerInstructions") or []:
        yield from inner["instructions"]


def analyse(tx):
    meta = tx["meta"]
    keys = [k["pubkey"] for k in tx["transaction"]["message"]["accountKeys"]]
    payer = keys[0]
    pre_tok, post_tok = token_rows(meta, "preTokenBalances"), token_rows(meta, "postTokenBalances")

    native = Decimal(meta["postBalances"][0] - meta["preBalances"][0]) / LAMPORTS
    sol_equiv, stable, rent, unpriced = native, Decimal(0), Decimal(0), collections.Counter()
    for idx in set(pre_tok) | set(post_tok):
        row = post_tok.get(idx) or pre_tok.get(idx)
        if row.get("owner") != payer:
            continue
        lamport_delta = Decimal(meta["postBalances"][idx] - meta["preBalances"][idx]) / LAMPORTS
        if row["mint"] == WSOL:
            sol_equiv += lamport_delta
            continue
        rent += lamport_delta
        sol_equiv += lamport_delta
        delta = (amount(post_tok[idx]) if idx in post_tok else 0) - (amount(pre_tok[idx]) if idx in pre_tok else 0)
        if row["mint"] in STABLES:
            stable += delta
        elif delta:
            unpriced[row["mint"]] += delta

    fee = Decimal(meta["fee"]) / LAMPORTS
    base = Decimal(BASE_FEE_PER_SIGNATURE * len(tx["transaction"]["signatures"])) / LAMPORTS
    tips, venues = Decimal(0), set()
    for ix in walk_instructions(tx):
        venue = VENUES.get(ix.get("programId"))
        if venue:
            venues.add(venue)
        parsed = ix.get("parsed")
        if isinstance(parsed, dict) and parsed.get("type") == "transfer" and ix.get("program") == "system":
            info = parsed["info"]
            if info.get("source") == payer and info.get("destination") in JITO_TIP_ACCOUNTS:
                tips += Decimal(info["lamports"]) / LAMPORTS

    return {
        "signature": tx["transaction"]["signatures"][0],
        "slot": tx["slot"],
        "blockTime": tx["blockTime"],
        "payer": payer,
        "success": meta["err"] is None,
        "solEquivalentDelta": sol_equiv,
        "stableDelta": stable,
        "rentDelta": rent,
        "unpriced": unpriced,
        "baseFee": base,
        "priorityFee": fee - base,
        "jitoTip": tips,
        "venues": sorted(venues),
    }


def main():
    frozen = json.load(open(f"{EVIDENCE}/frozen_sample.json"))
    path = f"{EVIDENCE}/transactions.jsonl"
    opener = (lambda: gzip.open(path + ".gz", "rt")) if not os.path.exists(path) else (lambda: open(path))
    rows = [analyse(json.loads(line)) for line in opener() if line.strip()]
    frozen["count"] = len(rows)

    # Price basis: median SOL/stable rate implied by this sample's own swaps that move both legs materially.
    implied = [abs(r["stableDelta"] / r["solEquivalentDelta"]) for r in rows
               if abs(r["stableDelta"]) > 50 and abs(r["solEquivalentDelta"]) > Decimal("0.1")]
    sol_usd = statistics.median(implied) if implied else None

    by_payer = collections.defaultdict(list)
    for r in rows:
        r["netUsd"] = r["stableDelta"] + r["solEquivalentDelta"] * sol_usd if sol_usd else None
        by_payer[r["payer"]].append(r)

    span_h = Decimal(max(r["blockTime"] for r in rows) - min(r["blockTime"] for r in rows)) / 3600
    summary = {
        "program": frozen["program"],
        "frozenAtUtc": frozen["frozenAtUtc"],
        "signaturesFrozen": frozen["count"],
        "transactionsAnalysed": len(rows),
        "windowHours": round(float(span_h), 2),
        "solUsdBasis": round(float(sol_usd), 2) if sol_usd else None,
        "solUsdBasisSource": f"median of {len(implied)} in-sample SOL/stable swaps",
        "wallets": [],
    }
    for payer, txs in sorted(by_payer.items(), key=lambda kv: -len(kv[1])):
        ok = [t for t in txs if t["success"]]
        nets = [t["netUsd"] for t in txs]
        unpriced = collections.Counter()
        for t in txs:
            unpriced.update(t["unpriced"])
        venues = collections.Counter(v for t in txs for v in t["venues"])
        summary["wallets"].append({
            "wallet": payer,
            "transactions": len(txs),
            "succeeded": len(ok),
            "failed": len(txs) - len(ok),
            "netUsdPricedLegs": round(float(sum(nets)), 4),
            "netUsdPerTransaction": round(float(sum(nets) / len(txs)), 5),
            "medianNetUsdSuccessful": round(float(statistics.median([t["netUsd"] for t in ok])), 5) if ok else None,
            "bestNetUsd": round(float(max(nets)), 4),
            "worstNetUsd": round(float(min(nets)), 4),
            "costsSol": {
                "baseFees": float(sum(t["baseFee"] for t in txs)),
                "priorityFees": float(sum(t["priorityFee"] for t in txs)),
                "jitoTips": float(sum(t["jitoTip"] for t in txs)),
                "failedTxFees": float(sum(t["baseFee"] + t["priorityFee"] for t in txs if not t["success"])),
            },
            "rentDeltaSol": float(sum(t["rentDelta"] for t in txs)),
            "unpricedTokenDeltas": {k: float(v) for k, v in unpriced.items() if v},
            "topVenues": venues.most_common(6),
        })

    all_nets = [r["netUsd"] for r in rows]
    summary["total"] = {
        "netUsdPricedLegs": round(float(sum(all_nets)), 4),
        "netUsdPerHour": round(float(sum(all_nets) / span_h), 4) if span_h else None,
        "successRatePct": round(100 * sum(r["success"] for r in rows) / len(rows), 1),
    }
    with open(f"{EVIDENCE}/per_transaction.csv", "w") as out:
        out.write("signature,slot,blockTime,payer,success,netUsd,stableDelta,solEquivalentDelta,baseFeeSol,priorityFeeSol,jitoTipSol,rentDeltaSol,unpricedTokens,venues\n")
        for r in rows:
            out.write(",".join(str(x) for x in [
                r["signature"], r["slot"], r["blockTime"], r["payer"], r["success"], r["netUsd"],
                r["stableDelta"], r["solEquivalentDelta"], r["baseFee"], r["priorityFee"], r["jitoTip"],
                r["rentDelta"], len(r["unpriced"]), "|".join(r["venues"]),
            ]) + "\n")
    json.dump(summary, open(f"{EVIDENCE}/summary.json", "w"), indent=2)
    print(json.dumps(summary, indent=2))


if __name__ == "__main__":
    main()
