from pathlib import Path
import json
import numpy as np
import pandas as pd

R = Path(__file__).resolve().parent
ROOT = R.parent

def load():
    g = pd.read_parquet(ROOT/'analysis-2026-09-05/games.parquet')
    c = pd.read_parquet(ROOT/'build-expression-2026-09-05/classification.parquet')
    m = pd.read_parquet(ROOT/'build-expression-2026-09-05/analysis_membership.parquet')
    d = g.merge(c, on='appid', validate='one_to_one').merge(
        m[['appid','paid','eligible','build','date_bin','quantity_max']], on='appid', validate='one_to_one')
    t = pd.read_parquet(ROOT/'analysis-2026-09-05/tags.parquet')
    d = d.set_index('appid', drop=False)
    for col, tag in [('deck','Roguelike Deckbuilder'),('action','Action Roguelike'),
                     ('auto','Auto Battler'),('tower','Tower Defense')]:
        d[col] = d.index.isin(t.loc[t.tag.eq(tag)&t['rank'].le(20),'appid'])
    d['text_words'] = (d.short_description.fillna('')+' '+d.description.fillna('')).str.split().str.len()
    d['text_bucket'] = pd.cut(d.text_words, [-1,150,300,600,np.inf],labels=['<=150','151-300','301-600','>600']).astype(str)
    d['period'] = np.select([x.fillna(False).to_numpy(dtype=bool) for x in [d.year.eq(2025),d.year.eq(2026)&d.month.le(6),d.year.eq(2026)&d.month.between(7,8)]],['2025','2026_H1','2026_JulAug'],default='other')
    d['pilot_family'] = np.select([d.deck,d.auto,d.action,d.tower],['deck','auto','action','tower'],default='other')
    return d,t

def stats(x):
    r=x.reviews
    n=len(x)
    a=x[r.ge(100)]
    return dict(n=n, under10=int(r.lt(10).sum()), ge10=int(r.ge(10).sum()),
                ge50=int(r.ge(50).sum()), ge100=int(r.ge(100).sum()),ge200=int(r.ge(200).sum()),ge1000=int(r.ge(1000).sum()),
                liked100=int((r.ge(100)&x.positive_pct.ge(80)).sum()),
                rate100=float(r.ge(100).mean()) if n else None,
                liked100_rate=float((r.ge(100)&x.positive_pct.ge(80)).mean()) if n else None,
                satisfied_among100=float(a.positive_pct.ge(80).mean()) if len(a) else None,
                median_reviews=float(r.median()) if n else None,
                median_positive_among100=float(a.positive_pct.median()) if len(a) else None,
                median_price=float(x.usd_list_price.median()) if x.usd_list_price.notna().any() else None,
                developers=int(x.developer_key.nunique()),publishers=int(x.publisher_key.nunique()),
                reviews_sum=int(r.sum()),top1_review_share=float(r.max()/r.sum()) if r.sum() else None)

def expected(pool, outcome, cols=('date_bin','price_bucket'), minimum=20):
    """Reproduce prior calibrated, full-pool standardization; not causal matching."""
    y=outcome.astype(float)
    date=pool.date_bin.astype(str)
    p=y.groupby(date).transform('mean')
    for n in range(2,len(cols)+1):
        key=pool[list(cols[:n])].fillna('missing').astype(str).agg('|'.join,axis=1)
        p=y.groupby(key).transform('mean').where(y.groupby(key).transform('size').ge(minimum),p)
    target=y.groupby(date).transform('sum')
    for _ in range(30):
        total=p.groupby(date).transform('sum')
        p=(p*target.div(total.where(total.gt(0))).fillna(0)).clip(upper=1)
        if (p.groupby(date).transform('sum')-target).abs().max()<1e-8: break
    assert abs(p.sum()-y.sum())<1e-5
    return p

def save_json(name,obj):
    (R/name).write_text(json.dumps(obj,indent=2,ensure_ascii=False,default=str))
