"""Reproduce the ProteinIQ ClinicalTrials.gov audit with Python 3.11+ (stdlib).

Fetch a new, version-checked snapshot: python analyze.py --fetch
Reproduce frozen results: python analyze.py --output-dir RESULTS
Labels in reviewed.tsv and review metadata in review.json are authored separately,
not inferred by this script. Keep all input files beside this script.
"""
import argparse
import csv
import gzip
import hashlib
import io
import json
import math
import platform
import random
import time
import urllib.parse
import urllib.request
from collections import Counter
from datetime import datetime, timezone
from pathlib import Path

ROOT = Path(__file__).resolve().parent
API = 'https://clinicaltrials.gov/api/v2/'
QUERY = ('AREA[StudyType]INTERVENTIONAL AND '
         '(AREA[InterventionType]DRUG OR AREA[InterventionType]BIOLOGICAL) AND '
         'AREA[StartDate]RANGE[2015-01-01,2020-12-31]')
FIELDS = ('NCTId,BriefTitle,StudyType,OverallStatus,WhyStopped,Phase,StartDate,StartDateType,'
          'EnrollmentCount,EnrollmentType,LeadSponsorClass,LastUpdatePostDate,'
          'StatusVerifiedDate,InterventionType,Condition,StudyFirstPostDate')
PHASES = {('PHASE1',): 'Phase 1', ('PHASE1','PHASE2'): 'Phase 1/2',
          ('PHASE2',): 'Phase 2', ('PHASE2','PHASE3'): 'Phase 2/3', ('PHASE3',): 'Phase 3'}
MAIN_SEED, PILOT_SEED, AUDIT_SEED = 20260919, 20260918, 20260920

def write_json(path, value):
    path.write_text(json.dumps(value, ensure_ascii=False, indent=2)+'\n', encoding='utf-8')

def digest(data):
    return hashlib.sha256(data).hexdigest()

def save_snapshot_parts(raw_pages):
    files = []
    for offset in range(0,len(raw_pages),20):
        name = f'snapshot-{offset//20+1:02}.json.gz'
        payload = gzip.compress(json.dumps(raw_pages[offset:offset+20],ensure_ascii=False).encode('utf-8'),mtime=0)
        (ROOT/name).write_bytes(payload)
        files.append({'file':name,'sha256':digest(payload),'pages':len(raw_pages[offset:offset+20])})
    return files

def request(url):
    for attempt in range(4):
        try:
            with urllib.request.urlopen(urllib.request.Request(url, headers={
                'User-Agent': 'ProteinIQ-reproducible-research/1.0',
                'Accept': 'application/json'}), timeout=90) as response:
                return response.read()
        except (OSError, TimeoutError):
            if attempt == 3:
                raise
            time.sleep(2 ** attempt)

def fetch_snapshot():
    version_before = json.loads(request(API+'version'))
    started = datetime.now(timezone.utc).isoformat()
    raw_pages, manifest, token, count = [], [], None, 0
    seen_tokens = set()
    while True:
        params = {'format':'json', 'query.term':QUERY, 'fields':FIELDS,
                  'pageSize':1000, 'countTotal':'true'}
        if token:
            params['pageToken'] = token
        url = API+'studies?'+urllib.parse.urlencode(params)
        raw = request(url)
        page = json.loads(raw)
        count += len(page['studies'])
        raw_pages.append(raw.decode('utf-8'))
        manifest.append({'url':url, 'response_sha256':digest(raw),
                         'records':len(page['studies']), 'total':page.get('totalCount')})
        print(f'Page {len(raw_pages)}: {count}/{page.get("totalCount")}', flush=True)
        token = page.get('nextPageToken')
        if not token:
            break
        if token in seen_tokens:
            raise ValueError('Repeated pagination token')
        seen_tokens.add(token)
    version_after = json.loads(request(API+'version'))
    if version_before != version_after:
        raise ValueError('API data version changed during retrieval; rerun extraction')
    studies = [s for raw in raw_pages for s in json.loads(raw)['studies']]
    ids = [s['protocolSection']['identificationModule']['nctId'] for s in studies]
    if (len(set(ids)) != count or manifest[0]['total'] != count
            or any(p['total'] not in (None, count) for p in manifest)):
        raise ValueError('Duplicate IDs or pagination/count mismatch')
    snapshot_files = save_snapshot_parts(raw_pages)
    write_json(ROOT/'sources.json', {'api':API, 'query':QUERY, 'fields':FIELDS,
        'retrieval_started':started, 'retrieval_finished':datetime.now(timezone.utc).isoformat(),
        'version_before':version_before, 'version_after':version_after,
        'python':platform.python_version(), 'records':count,
        'snapshot_files':snapshot_files, 'pages':manifest})

def load_snapshot():
    provenance = json.loads((ROOT/'sources.json').read_text(encoding='utf-8'))
    pages = []
    for part in provenance['snapshot_files']:
        raw = (ROOT/part['file']).read_bytes()
        if digest(raw) != part['sha256']:
            raise ValueError('Snapshot checksum mismatch')
        chunk = json.loads(gzip.decompress(raw))
        assert len(chunk)==part['pages']
        pages.extend(chunk)
    assert len(pages) == len(provenance['pages'])
    studies = []
    for body, page in zip(pages, provenance['pages']):
        assert digest(body.encode('utf-8')) == page['response_sha256']
        parsed = json.loads(body)
        assert len(parsed['studies']) == page['records']
        studies.extend(parsed['studies'])
    assert len(studies) == provenance['records']
    ids = [s['protocolSection']['identificationModule']['nctId'] for s in studies]
    assert len(set(ids)) == len(ids)
    return studies, provenance

def flatten(study):
    p = study['protocolSection']
    status, design = p.get('statusModule',{}), p.get('designModule',{})
    start, enroll = status.get('startDateStruct',{}), design.get('enrollmentInfo',{})
    phases = tuple(sorted(design.get('phases',[])))
    interventions = sorted({x['type'] for x in p.get('armsInterventionsModule',{}).get('interventions',[])})
    return {'nct_id':p['identificationModule']['nctId'],
        'title':p['identificationModule'].get('briefTitle',''),
        'study_type':design.get('studyType',''),
        'phase':PHASES.get(phases,''), 'phase_raw':'|'.join(phases),
        'intervention_types':'|'.join(interventions),
        'start_date':start.get('date',''), 'start_type':start.get('type',''),
        'status':status.get('overallStatus',''), 'why_stopped':status.get('whyStopped',''),
        'enrollment':enroll.get('count'), 'enrollment_type':enroll.get('type',''),
        'sponsor_class':p.get('sponsorCollaboratorsModule',{}).get('leadSponsor',{}).get('class',''),
        'last_update':status.get('lastUpdatePostDateStruct',{}).get('date',''),
        'status_verified':status.get('statusVerifiedDate',''),
        'first_posted':status.get('studyFirstPostDateStruct',{}).get('date','')}

def exclusion(row):
    # First matching exclusion is used so counts are mutually exclusive.
    if row['study_type'] != 'INTERVENTIONAL': return 'not_interventional'
    if not set(row['intervention_types'].split('|')) & {'DRUG','BIOLOGICAL'}: return 'not_drug_or_biological'
    if not row['phase']: return 'phase_outside_scope'
    if row['start_type'] != 'ACTUAL': return 'start_not_actual'
    if not '2015' <= row['start_date'][:4] <= '2020': return 'start_outside_period'
    if row['status'] in {'WITHDRAWN','NOT_YET_RECRUITING'}: return 'start_status_conflict'
    if row['enrollment_type'] == 'ACTUAL' and row['enrollment'] == 0: return 'zero_actual_enrollment'
    if row['status'] not in {'COMPLETED','TERMINATED','SUSPENDED','RECRUITING',
        'ENROLLING_BY_INVITATION','ACTIVE_NOT_RECRUITING','UNKNOWN'}: return 'unsupported_status'
    return ''

def sample(rows, size, seed):
    ordered = sorted(rows, key=lambda r:r['nct_id'])
    return sorted(random.Random(seed).sample(ordered, min(size,len(ordered))), key=lambda r:r['nct_id'])

def wilson(k, n):
    # Marginal 95% Wilson interval, without finite-population correction.
    if not n: return None
    z, p = 1.959963984540054, k/n
    center = (p+z*z/(2*n))/(1+z*z/n)
    half = z*math.sqrt(p*(1-p)/n+z*z/(4*n*n))/(1+z*z/n)
    return [0.0 if k == 0 else 100*(center-half),
            100.0 if k == n else 100*(center+half)]

def stats(rows):
    counts = Counter(r['status'] for r in rows)
    return {'n':len(rows), 'terminated':counts['TERMINATED'],
        'terminated_pct':100*counts['TERMINATED']/len(rows) if rows else None,
        'statuses':dict(sorted(counts.items()))}

def csv_bytes(rows):
    out = io.StringIO(newline='')
    if rows:
        writer = csv.DictWriter(out, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)
    return out.getvalue().encode('utf-8')

def analyze(output):
    output.mkdir(parents=True,exist_ok=True)
    studies, provenance = load_snapshot()
    rows = sorted([flatten(s) for s in studies],key=lambda r:r['nct_id'])
    for row in rows: row['exclusion'] = exclusion(row)
    eligible = [r for r in rows if not r['exclusion']]
    terminated = [r for r in eligible if r['status']=='TERMINATED']
    main, pilot = sample(terminated,400,MAIN_SEED), sample(terminated,30,PILOT_SEED)
    main_ids, pilot_ids = {r['nct_id'] for r in main}, {r['nct_id'] for r in pilot}
    audit_ids = [r['nct_id'] for r in sample(main,40,AUDIT_SEED)]
    summary = {'retrieved_at':provenance['retrieval_finished'], 'data_version':provenance['version_after'],
        'retrieved_records':len(rows), 'exclusions':dict(Counter(r['exclusion'] for r in rows if r['exclusion'])),
        'overall':stats(eligible), 'phases':{phase:stats([r for r in eligible if r['phase']==phase]) for phase in PHASES.values()},
        'older_2015_2019':stats([r for r in eligible if r['start_date'][:4]<'2020']),
        'sponsors':{group:stats([r for r in eligible if (r['sponsor_class']=='INDUSTRY')==is_industry])
                    for group,is_industry in [('industry_led',True),('other_or_missing_lead',False)]},
        'missing_enrollment':sum(r['enrollment'] is None for r in eligible),
        'sample':{'n':len(main), 'seed':MAIN_SEED, 'pilot_seed':PILOT_SEED,'pilot_n':len(pilot),
                  'pilot_overlap':sorted(main_ids & pilot_ids), 'audit_seed':AUDIT_SEED, 'audit_ids':audit_ids}}
    assert summary['overall']['n']+sum(summary['exclusions'].values())==len(rows)
    for name,data in [('cohort.csv.gz',rows),('reason-sample.csv',main),('pilot.csv',pilot)]:
        payload = csv_bytes(data)
        (output/name).write_bytes(gzip.compress(payload,mtime=0) if name.endswith('.gz') else payload)
    review_path = ROOT/'review.json'
    if review_path.exists():
        review = json.loads(review_path.read_text(encoding='utf-8'))
        assert review['snapshot_sha256'] == {
            part['file']:part['sha256'] for part in provenance['snapshot_files']
        }, 'Review belongs to a different snapshot; review the new source text first'
        with (ROOT/'reviewed.tsv').open(encoding='utf-8',newline='') as handle:
            entries = list(csv.DictReader(handle,delimiter='\t'))
        assert all(r['uncertain'] in {'true','false'} for r in entries)
        decisions = {r['nct_id']:{'labels':r['labels'].split('|'),
            'uncertain':r['uncertain']=='true','rationale':r['rationale']} for r in entries}
        assert len(decisions)==len(entries), 'Duplicate review ID'
        assert set(decisions)==main_ids, 'Review must cover the exact sample'
        codebook = json.loads((ROOT/'codebook.json').read_text(encoding='utf-8'))
        counts = Counter()
        coded_rows = []
        for row in main:
            decision = decisions[row['nct_id']]
            labels = decision['labels']
            assert labels and len(set(labels))==len(labels)
            assert set(labels)<=set(codebook['labels'])
            assert decision['rationale'].strip()
            counts.update(labels)
            coded_rows.append({**row,'labels':'|'.join(labels),'uncertain':decision['uncertain'],
                               'rationale':decision['rationale']})
        assert set(review['rechecked_ids']) >= set(audit_ids)
        assert set(review['rechecked_ids']) >= {k for k,v in decisions.items() if v['uncertain']}
        assert set(review['rechecked_ids']) <= main_ids
        summary['reasons'] = {key:{'n':counts[key], 'pct':100*counts[key]/len(main),
            'wilson_95_pct':wilson(counts[key],len(main))} for key in codebook['labels']}
        summary['review'] = {'uncertain':sum(d['uncertain'] for d in decisions.values()),
            'rechecked':len(set(review['rechecked_ids'])), 'method':review['method']}
        conflicts = {k for k,v in decisions.items() if 'status_conflict' in v['labels']}
        summary['excluding_sample_flagged_conflicts'] = stats([r for r in eligible if r['nct_id'] not in conflicts])
        (output/'reason-sample.csv').write_bytes(csv_bytes(coded_rows))
    write_json(output/'summary.json',summary)
    print(json.dumps(summary,indent=2))

if __name__=='__main__':
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--fetch',action='store_true')
    parser.add_argument('--output-dir',type=Path,default=ROOT)
    args = parser.parse_args()
    if args.fetch: fetch_snapshot()
    analyze(args.output_dir)
