"""Phase 1 research arithmetic; not application code or an empirical forecast.

All non-source inputs are author assumptions. See phase-1-blueprint.md.
Run: python3 research/calculate_scenarios.py
"""
import csv
import json
import math
from pathlib import Path

ROOT = Path(__file__).parent
U0 = 0.038
FK0, FN0 = 62.4, 37.6
EK0, EN0 = FK0 * (1-U0), FN0 * (1-U0)
WAGE_RATIO = 1.5
BASE_BILL = WAGE_RATIO * EK0 + EN0
WK0, WN0 = WAGE_RATIO * 60/BASE_BILL, 60/BASE_BILL
THETA_K = WAGE_RATIO * EK0/BASE_BILL
ROOF_WEIGHTS = [.10, .12, .14, .12, .12, .22, .12, .06]
TASK_NAMES = ['access_safety', 'materials', 'tear_off', 'deck_repair',
              'underlayment', 'field_covering', 'flashing', 'handover']
PM_NAMES = ['inspection_drone', 'scope_estimates', 'purchasing', 'scheduling',
            'records_costs', 'site_safety_quality', 'people_coordination', 'exceptions']
PM_WEIGHTS = [.12,.16,.12,.15,.12,.15,.12,.06]
PM_DIGITAL = {
    'mild': [.05,.20,.15,.15,.25,0,.05,0],
    'moderate': [.15,.50,.50,.40,.70,.02,.10,.02],
    'aggressive': [.35,.80,.85,.75,.95,.10,.30,.05],
}
PARAMS = {
 'mild': dict(m=.2, diffusion=.2, auto=.5, reinstate=.5, log_gain=.3,
              target=.016, match=.7, train=.005, entrants=.03,
              robot=[0,.05,.02,0,.02,.04,0,0], augment=.10, gain=1.20),
 'moderate': dict(m=.3, diffusion=.4, auto=.75, reinstate=.25, log_gain=.45,
                  target=.083, match=.5, train=.01, entrants=.10,
                  robot=[0,.30,.10,.02,.15,.28,.01,.03], augment=.20, gain=1.25),
 'aggressive': dict(m=.5, diffusion=.6, auto=.9, reinstate=0, log_gain=.8,
                    target=.324, match=.35, train=.02, entrants=.20,
                    robot=[0,.80,.45,.15,.50,.75,.10,.15], augment=.30, gain=1.35),
}

def clip(x):
    return max(0., min(1., x))

def run(name, robotics=True, entrants=True, robot_start=2, demand_scale=1.,
        roof_volume_multiplier=1., wage_exponent=.5, physical_eligibility=.5):
    p = PARAMS[name]
    ek, en, uk, un, training = EK0, EN0, FK0-EK0, FN0-EN0, 0.
    wk, wn, wr = 1., 1., 1.
    rows = []
    robot_terminal = sum(w*a for w,a in zip(ROOF_WEIGHTS,p['robot']))
    for t in range(7):
        x = t/6
        # Shifting the start delays rather than accelerates the four-year ramp.
        r = clip((t-robot_start)/4) if robotics else 0.
        ak = p['m'] * p['diffusion'] * p['auto'] / .624 * x
        zk = p['m'] * p['diffusion'] * (1-p['auto']) / .624 * x
        hk = 1-ak-zk+zk/math.exp(p['log_gain'])+.10*ak+p['reinstate']*ak
        ar = robot_terminal*r
        an = physical_eligibility*ar
        zn = (1-an)*p['augment']*.5*x
        hn = 1-an-zn+zn/p['gain']+.15*an+p['reinstate']*an
        q = (1+p['target']*x)*(1+.20*an)*(1+(demand_scale-1)*x)
        dk, dn = EK0*q*hk, EN0*q*hn
        starts = completions = 0.
        if t:
            completions, training = training, 0.
            un += completions
            lostk, lostn = max(0.,ek-dk), max(0.,en-dn)
            ek -= lostk; uk += lostk
            en -= lostn; un += lostn
            hirek = min(max(0., dk-ek), p['match']*uk)
            hiren = min(max(0., dn-en), p['match']*un)
            ek += hirek; uk -= hirek
            en += hiren; un -= hiren
            # Transfer only the excess K unemployment into a 1-year pipeline.
            # All trainees stay active jobseekers in this toy model.
            if entrants:
                starts = min(max(0., uk-U0*(ek+uk))*.5, p['train']*FK0)
                uk -= starts; training += starts
            sk, sn = ek+uk+training, en+un
            # Bargaining rule conditional on fixed order volume, not equilibrium.
            targetk = ((dk/sk)/(1-U0))**wage_exponent
            targetn = ((dn/sn)/(1-U0))**wage_exponent
            wk = math.sqrt(wk*targetk)
            wn = math.sqrt(wn*targetn)
        else:
            sk, sn = FK0,FN0
        qk, qn = ek/(EK0*hk), en/(EN0*hn)
        gdp = 100*(THETA_K*qk+(1-THETA_K)*qn)
        bill = WK0*wk*ek+WN0*wn*en
        residual = gdp-bill
        # Independent roofing receiving-market sensitivity, not a slice of
        # national employment stocks. Entrant path is net qualified supply.
        hroof = 0.
        for weight,terminal in zip(ROOF_WEIGHTS,p['robot']):
            a = terminal*r
            z = (1-a)*p['augment']*x
            assert a+z <= 1+1e-12
            hroof += weight*(1-a-z+z/p['gain']+.15*a+p['reinstate']*a)
        inflow = p['entrants']*clip((t-1)/5) if entrants else 0.
        volume = (gdp/100)**.5*hroof**(-.2)*(1+(roof_volume_multiplier-1)*x)
        jobs = volume*hroof
        roof_pool = 1/(1-U0)+inflow
        relative_roof_pool = roof_pool/(1/(1-U0))
        pressure = relative_roof_pool/jobs
        targetr = pressure**(-wage_exponent)
        if t: wr = math.sqrt(wr*targetr)
        desired_roof_emp = min(jobs,roof_pool)
        roof_u = 100*(1-desired_roof_emp/roof_pool)
        row = dict(scenario=name, model_year=t, gdp_index=gdp,
            knowledge_wage_index=100*wk, nonknowledge_wage_index=100*wn,
            roofing_wage_index=100*wr, unemployment_pct=uk+un+training,
            knowledge_jobseeker_unemployment_pct=100*(uk+training)/sk,
            nonknowledge_unemployment_pct=100*un/sn,
            jobseekers_excluding_training_per100=uk+un,
            training_per100=training, training_starts=starts,
            training_completions=completions, employed_k=ek, employed_n=en,
            unemployed_k=uk, unemployed_n=un,
            labor_share_pct=100*bill/gdp, labor_income_index=100*bill/60,
            capital_residual_index=100*residual/40,
            human_hours_k=hk, human_hours_n=hn, human_hours_roof=hroof,
            roof_robot_task_share_pct=100*ar, roof_entrant_stock_pct=100*inflow,
            roof_entrant_stock_scale=166900*inflow,
            roof_output_index=100*volume, roof_job_demand_index=100*jobs,
            roof_output_needed_for_baseline_tightness_index=100*relative_roof_pool/hroof,
            roof_pressure_index=100*pressure, roof_static_capacity_gap_pct=roof_u,
            target_output_index=100*q,
            labor_bill_units=bill, capital_residual_units=residual,
        )
        assert abs(ek+en+uk+un+training-100) < 1e-9
        assert min(ek,en,uk,un,training,hk,hn,hroof,bill,residual)>=-1e-9
        assert abs(gdp-(bill+residual))<1e-9
        assert 0<=row['labor_share_pct']<=100
        assert ak+zk<=1
        rows.append(row)
    return rows

def main():
    rows=[row for name in PARAMS for row in run(name)]
    assert abs(sum(ROOF_WEIGHTS)-1)<1e-12
    assert abs(sum(PM_WEIGHTS)-1)<1e-12
    for name in PARAMS:
        base=run(name)[0]
        assert abs(base['gdp_index']-100)<1e-9
        assert abs(base['labor_share_pct']-60)<1e-9
        assert abs(base['unemployment_pct']-3.8)<1e-9
    with (ROOT/'scenario-data.csv').open('w') as f:
        w=csv.DictWriter(f,fieldnames=rows[0].keys(),lineterminator='\n');w.writeheader();w.writerows(rows)
    payload=dict(schema_version='0.1-draft', classification='author-constructed conditional stress test',
        baseline='constant no-additional-automation counterfactual = 100; no calendar date',
        notice='Not an empirical forecast, not a replication of Anthropic, no scenario probabilities.',
        parameters=PARAMS,global_assumptions=dict(baseline_unemployment=U0,knowledge_labor_force_share=.624,initial_knowledge_wage_ratio=WAGE_RATIO,baseline_labor_share=.6,macro_physical_eligibility=.5,independent_nonknowledge_augmentation_access=.5,digital_support_hours_per_automated_hour=.10,robot_support_hours_per_automated_hour=.15,wage_pressure_exponent=.5,annual_log_wage_adjustment=.5,training_years=1,baseline_growth=0,roofing_income_elasticity=.5,roofing_price_elasticity_times_cost_pass_through=.2,robot_output_demand_uplift=.20),roof_task_names=TASK_NAMES,roof_task_weights=ROOF_WEIGHTS,rows=rows)
    (ROOT/'scenario-data.json').write_text(json.dumps(payload,indent=2)+'\n')
    tasks = dict(classification='author-assigned time weights and terminal deployment shares; not measured O*NET times', instance_share_basis='weighted by baseline human duration within each task bundle, not unweighted instance counts', roles={})
    for role,names,weights in [('production_manager',PM_NAMES,PM_WEIGHTS),('roofer',TASK_NAMES,ROOF_WEIGHTS)]:
        entries=[]
        for i,(task,weight) in enumerate(zip(names,weights)):
            states={}
            for scenario,p in PARAMS.items():
                a=PM_DIGITAL[scenario][i] if role=='production_manager' else p['robot'][i]
                z=(1-a)*p['augment']
                states[scenario]=dict(automated_baseline_instances=a, augmented_baseline_instances=z,
                    unassisted_baseline_instances=1-a-z,augmentation_productivity=p['gain'],
                    human_support_hours_per_automated_baseline_hour=.10 if role=='production_manager' else .15,
                    new_human_hours_per_automated_baseline_hour=p['reinstate'])
                assert abs(a+z+(1-a-z)-1)<1e-12
            entries.append(dict(id=task,baseline_time_weight=weight,scenarios=states))
        tasks['roles'][role]=entries
    (ROOT/'task-bundles.json').write_text(json.dumps(tasks,indent=2)+'\n')
    sensitivity=[]
    for name in PARAMS:
        for label,kwargs in [('full',{}),('no_direct_robot_substitution',dict(robotics=False)),
                 ('no_entrants',dict(entrants=False)),
                 ('no_direct_robot_substitution_or_entrants',dict(robotics=False,entrants=False)),
                 ('robotics_delayed_2y',dict(robot_start=4)),
                 ('demand_10pct_lower',dict(demand_scale=.9)),
                 ('roof_demand_25pct_higher',dict(roof_volume_multiplier=1.25)),
                 ('wage_response_weak',dict(wage_exponent=.25)),
                 ('wage_response_strong',dict(wage_exponent=.75)),
                 ('physical_eligibility_low',dict(physical_eligibility=.25)),
                 ('physical_eligibility_high',dict(physical_eligibility=.75))]:
            sensitivity.append(dict(case=label,**run(name,**kwargs)[-1]))
    (ROOT/'scenario-sensitivity.json').write_text(json.dumps(sensitivity,indent=2)+'\n')
    for r in rows:
        if r['model_year'] in (0,2,4,6):
            print(r['scenario'],r['model_year'], 'GDP %.2f K %.2f N %.2f Roof %.2f U %.2f LS %.2f' % tuple(r[k] for k in ['gdp_index','knowledge_wage_index','nonknowledge_wage_index','roofing_wage_index','unemployment_pct','labor_share_pct']))
    print('Verified: 21 annual rows; population conservation; task partitions; baseline; income identity; 33 sensitivity endpoints.')

if __name__=='__main__': main()
