"""Deterministic Lab 14 accounting; artificial units, no model or external tools."""
import argparse
import csv
import hashlib
import json
import math
from pathlib import Path

GOLD = frozenset({'C1', 'C2', 'C3'})  # Evaluation only; never passed to ranking/packing.
BODY_UNITS = dict(zip(('C1','C2','C3','C4','C5','C6','P','T','H','Q','H_short','H_faithful','S'),
                      (27,31,32,32,27,28,27,12,16,28,7,7,20)))
CHUNK_DOCUMENTS = dict(zip(('C1','C2','C3','C4','C5','C6'), ('D1','D1','D1','D2','D3','D3')))


def validate_fixture(fixture):
    """Reject malformed supplied evidence before ranking, packing or writing."""
    def keys(value, expected):
        if type(value) is not dict or set(value) != set(expected):
            raise ValueError('Invalid fixture fields or IDs')
    keys(fixture, ('bodies','initial','reranked','sources','chunks'))
    keys(fixture['bodies'], BODY_UNITS)
    for key, units in BODY_UNITS.items():
        body = fixture['bodies'][key]
        if type(body) is not str or len(body.split()) != units:
            raise ValueError('Invalid fixture body or unit count: '+key)
    keys(fixture['sources'], ('D1','D2','D3'))
    for source in fixture['sources'].values():
        keys(source, ('version','status'))
        if any(type(value) is not str or not value.strip() for value in source.values()):
            raise ValueError('Invalid source metadata')
    keys(fixture['chunks'], CHUNK_DOCUMENTS)
    for key, document in CHUNK_DOCUMENTS.items():
        chunk = fixture['chunks'][key]
        keys(chunk, ('document','section','instruction_authority','trust'))
        if (chunk['document'] != document or chunk['document'] not in fixture['sources']
                or type(chunk['section']) is not str or not chunk['section'].strip()
                or chunk['instruction_authority'] != 'none' or chunk['trust'] != 'source_data'):
            raise ValueError('Invalid chunk provenance: '+key)
    for name in ('initial','reranked'):
        keys(fixture[name], CHUNK_DOCUMENTS)
        if any(type(score) not in (int,float) or not math.isfinite(score)
               for score in fixture[name].values()):
            raise ValueError('Invalid ranking scores')


def rank(scores):
    return sorted(scores, key=lambda key: (-scores[key], key))


def pack(order, bodies, allowance):
    remaining, included, rejected = allowance, [], []
    for key in order:
        cost = len(bodies[key].split()) + 8
        if cost <= remaining:
            included.append(key)
            remaining -= cost
        else:
            rejected.append({'id': key, 'cost': cost, 'reason': 'whole_record_exceeds_remaining'})
    return {'included': included, 'rejected': rejected, 'evidence_cost': allowance-remaining,
            'unused_evidence_allowance': remaining}


def evaluate(ids, gold=GOLD):
    selected = set(ids)
    hits = len(selected & gold)
    return {'precision': hits/len(selected) if selected else None, 'recall': hits/len(gold)}


def replay(fixture):
    validate_fixture(fixture)
    bodies = fixture['bodies']
    costs = {key: len(value.split()) + 8 for key, value in bodies.items()}
    mandatory = sum(costs[key] for key in ('P','T','H','Q'))
    allowance = 300-60-11-mandatory
    variants = {}
    def variant(name, order, budget, subtotal):
        result = pack(order, bodies, budget)
        result.update(evaluation=evaluate(result['included']), total_input_cost=subtotal+result['evidence_cost'],
                      mandatory_cost=subtotal, allowance=budget, ranking=order)
        variants[name] = result
    initial, reranked = rank(fixture['initial']), rank(fixture['reranked'])
    variant('A', initial, allowance, mandatory)
    variant('B', reranked, allowance, mandatory)
    for history in ('H_short','H_faithful'):
        subtotal=mandatory-costs['H']+costs[history]
        variant(history, reranked, 300-60-11-subtotal, subtotal)
    retrieval = {'initial': {str(k):evaluate(initial[:k]) for k in (2,4,5)},
                 'reranked_top3': evaluate(reranked[:3])}
    return {'kind':'deterministic_fixture_replay', 'model_inference':'not_run',
            'external_tool_execution':'not_run', 'unit':'whitespace_body_plus_eight_header',
            'scores':'supplied_fixture_not_model_measurements', 'costs': costs,
            'sources':fixture['sources'], 'chunks':fixture['chunks'], 'retrieval':retrieval,
            'variants':variants, 'summary':{'references':['C1','C2','C3'], 'cost':costs['S'],
                                          'total_input_cost':mandatory+costs['S']}}


def fresh_output(path):
    path=Path(path).absolute()
    for ancestor in [path,*path.parents]:
        if ancestor.is_symlink(): raise ValueError('Output path contains a symlink')
    path.mkdir(mode=0o700,parents=True,exist_ok=False)
    return path


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--fixture',type=Path,default=Path(__file__).with_name('product_trace_fixture.json'))
    parser.add_argument('--output',required=True,type=Path)
    args=parser.parse_args()
    fixture=json.loads(args.fixture.read_text())
    trace=replay(fixture)
    trace['fixture_sha256']=hashlib.sha256(args.fixture.read_bytes()).hexdigest()
    trace['runner_sha256']=hashlib.sha256(Path(__file__).read_bytes()).hexdigest()
    checks={'mandatory_115':trace['variants']['A']['mandatory_cost']==115,
            'A_ids':trace['variants']['A']['included']==['C1','C4','C2'],
            'B_ids':trace['variants']['B']['included']==['C1','C3','C2'],
            'summary_143':trace['summary']['total_input_cost']==143,
            'zero_allowance':pack(rank(fixture['initial']),fixture['bodies'],0)['included']==[],
            '34_allowance':pack(rank(fixture['initial']),fixture['bodies'],34)['included']==[],
            'compression_unchanged':all(trace['variants'][k]['included']==['C1','C3','C2'] for k in ('H_short','H_faithful')),
            'instruction_is_source_data':fixture['chunks']['C6']['instruction_authority']=='none'}
    if not all(checks.values()): raise ValueError('Fixture check failed')
    output=fresh_output(args.output)
    (output/'trace.json').write_text(json.dumps(trace,indent=2)+'\n')
    (output/'checks.json').write_text(json.dumps(checks,indent=2)+'\n')
    with (output/'packing.csv').open('w',newline='') as handle:
        writer=csv.DictWriter(handle,fieldnames=['variant','included','evidence_cost','total_input_cost','unused_evidence_allowance','precision','recall'])
        writer.writeheader()
        for name,result in trace['variants'].items():
            writer.writerow({'variant':name,'included':' '.join(result['included']),
                             **{k:result[k] for k in ('evidence_cost','total_input_cost','unused_evidence_allowance')},
                             **result['evaluation']})
    print('Deterministic fixture checks passed; model inference not run.')


if __name__=='__main__': main()
