#!/usr/bin/env python3
"""Finish sampled buckets — operator GO 2026-08-18.

Previous run died mid-listing of openobserve (1.7M objs drained the 1h token).
This run:
  - per-bucket token re-mint (no mid-bucket 401)
  - FULL download (not sample) of 9 buckets:
      cfu-main-openkm, cfu-main-data-warehouse, cfu-main-factory-patroli-image,
      cfu-main.appspot.com, storage-innopharm-prod, ppds, cfu-main-omniagent,
      cfu-main-file-kontrol-perubahan, staging.cfu-main.appspot.com
  - openobserve: SKIP bulk (RUM log spam, 1.7M objs); the 30-sample from prev run
    is enough intel. We just record the skip.
  - skips existing files by size
Egress: storage.googleapis.com + oauth2.googleapis.com only.
Out: downloads/all/<bucket>/ + sampled_finish_index_aug18.json + OPLOG.
"""
import base64
import importlib.util
import json
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(60)

ROOT = Path('/root/ir-assessment')
DOSSIER = ROOT / 'redteam/gitlab_pharmalink_id'
OUT = DOSSIER / 'downloads/all'
INDEX = DOSSIER / 'downloads/sampled_finish_index_aug18.json'

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

BUCKETS = [
    # 'cfu-main-openkm',          # DONE 2026-08-18 (288 real files, rest = folder markers)
    # 'cfu-main-data-warehouse',  # DONE 2026-08-18 (406 real files)
    'cfu-main-factory-patroli-image',
    'cfu-main.appspot.com',
    'storage-innopharm-prod',
    'ppds',
    'cfu-main-omniagent',
    'cfu-main-file-kontrol-perubahan',
    'staging.cfu-main.appspot.com',
]
PROJ = 'cfu-main'
# Files above this size get individually logged but still downloaded (cap safety)
HUGE_CAP = 2 * 1024**3  # 2GB per single object


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 (per-bucket re-mint, full dl) | finish sampled buckets | {msg} | {result} | none | sampled-finish 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 list_all(tok, bucket):
    items, page = [], None
    while True:
        url = f'https://storage.googleapis.com/storage/v1/b/{bucket}/o?maxResults=1000&fields=items(name,size),nextPageToken'
        if page:
            url += '&pageToken=' + urllib.parse.quote(page)
        try:
            req = urllib.request.Request(url, headers={'Authorization': f'Bearer {tok}',
                                                       'Connection': 'close'})
            with urllib.request.urlopen(req, timeout=25, context=CTX) as resp:
                body = resp.read().decode('utf-8', 'ignore')
            r = json.loads(body)
        except Exception as e:
            print(f'  [!] {bucket} list page err: {type(e).__name__} {e}', file=sys.stderr)
            return items if items else None
        items.extend((i['name'], int(i.get('size', 0))) for i in r.get('items', []))
        page = r.get('nextPageToken')
        if not page:
            return items


def dl(tok, bucket, name, dest):
    enc = urllib.parse.quote(name, safe='')
    url = f'https://storage.googleapis.com/download/storage/v1/b/{bucket}/o/{enc}?alt=media'
    r = urllib.request.Request(url, headers={'Authorization': f'Bearer {tok}'})
    with urllib.request.urlopen(r, timeout=300, context=CTX) as resp, open(dest, 'wb') as f:
        while True:
            chunk = resp.read(1 << 20)
            if not chunk:
                break
            f.write(chunk)
    return dest.stat().st_size


class TokenMgr:
    """Token lifecycle: re-mint every TOKEN_TTL sec OR on 401."""
    TTL = 50 * 60  # 50 min

    def __init__(self, sa):
        self.sa = sa
        self.token = None
        self.minted_at = 0

    def get(self):
        if not self.token or (time.time() - self.minted_at) > self.TTL:
            self.token = mint(self.sa)
            self.minted_at = time.time()
        return self.token

    def invalidate(self):
        self.token = None
        self.minted_at = 0


def dl_with_retry(tm, bucket, name, dest, retries=3):
    """dl with auto-retry on 401 (token expiry) — re-mint and retry."""
    last_err = None
    for attempt in range(retries):
        tok = tm.get()
        try:
            return dl(tok, bucket, name, dest)
        except urllib.error.HTTPError as e:
            last_err = e
            if e.code == 401 and attempt < retries - 1:
                tm.invalidate()
                time.sleep(2)
                continue
            raise
        except Exception as e:
            last_err = e
            if attempt < retries - 1:
                time.sleep(3)
                continue
            raise
    raise last_err


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

    index = []
    grand_ok = grand_skip = grand_fail = 0

    for bucket in BUCKETS:
        print(f'[*] bucket: {bucket} (fresh token)', file=sys.stderr)
        tm = TokenMgr(sa)  # per-bucket TokenMgr with in-bucket refresh
        tok = tm.get()

        items = list_all(tok, bucket)
        if items is None:
            index.append({'bucket': bucket, 'fatal': 'list_failed'})
            print(f'  [!] list failed', file=sys.stderr)
            continue

        bdir = OUT / bucket
        bdir.mkdir(parents=True, exist_ok=True)
        n_ok = n_skip = n_fail = n_huge = 0

        for name, size in items:
            if size == 0 and name.endswith('/'):
                continue
            if size > HUGE_CAP:
                index.append({'bucket': bucket, 'name': name, 'ok': False,
                              'skipped_huge': True, 'size': size})
                n_huge += 1
                print(f'  [skip-huge] {name[:60]} {size/1e9:.2f}GB', file=sys.stderr)
                continue
            slug = name.replace('/', '__')
            dest = bdir / slug
            if dest.exists() and dest.stat().st_size == size:
                n_skip += 1
                continue
            try:
                got = dl_with_retry(tm, bucket, name, dest)
                ok = got == size
                index.append({'bucket': bucket, 'name': name, 'bytes': got, 'ok': ok})
                if ok:
                    n_ok += 1
                else:
                    n_fail += 1
            except Exception as e:
                index.append({'bucket': bucket, 'name': name, 'ok': False,
                              'error': str(e)[:120]})
                n_fail += 1
                print(f'  [-] {name[:60]}: {e}', file=sys.stderr)

        index.append({'bucket': bucket, 'summary': True, 'total': len(items),
                      'ok': n_ok, 'skip': n_skip, 'fail': n_fail, 'huge': n_huge})
        grand_ok += n_ok; grand_skip += n_skip; grand_fail += n_fail
        print(f'  [{bucket}] total={len(items)} ok={n_ok} skip={n_skip} fail={n_fail} huge={n_huge}',
              file=sys.stderr)
        INDEX.write_text(json.dumps(index, indent=1))

    INDEX.write_text(json.dumps(index, indent=1))
    oplog(f'buckets={len(BUCKETS)} ok={grand_ok} skip={grand_skip} fail={grand_fail}',
          'SUCCESS' if grand_fail == 0 else 'PARTIAL')
    print(f'[+] grand: ok={grand_ok} skip={grand_skip} fail={grand_fail}', file=sys.stderr)


if __name__ == '__main__':
    main()
