#!/usr/bin/env python3
"""Local-only TAPNext++ adapter for jobs exported by the BowFlow browser MVP.

Uses the official TAPNext PyTorch model and the official 256px checkpoint.
Nothing is uploaded. No synthetic output or fallback model is used here.
"""
from __future__ import annotations

import argparse
import hashlib
import json
import math
from pathlib import Path
import subprocess
import sys
import time

ALLOWED_POINTS = {'bow_a', 'bow_b', 'bridge_a', 'bridge_b', 'string_a', 'string_b'}
UPSTREAM_COMMIT = '730cda1c730877cfedbe01bf87fb1cadb78a565d'
CHECKPOINT_SHA256 = 'cb96a43444ccb4fbdb25d800b88c7ba196179a526e78f01b021a16b1c1eff6da'


def sha256(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open('rb') as stream:
        for chunk in iter(lambda: stream.read(4 * 1024 * 1024), b''):
            digest.update(chunk)
    return digest.hexdigest()


def validate_job(job: dict) -> None:
    if job.get('schema') != 'bowflow-tapnext-job-v1' or job.get('engine') != 'tapnext++':
        raise ValueError('Not a BowFlow TAPNext++ job')
    if not isinstance(job.get('jobId'), str) or not job['jobId']:
        raise ValueError('Missing job ID')
    source = job.get('source', {})
    if len(source.get('sha256', '')) != 64 or any(c not in '0123456789abcdef' for c in source['sha256']):
        raise ValueError('Invalid source SHA-256')
    for name in ('width', 'height'):
        if not isinstance(source.get(name), int) or not 16 <= source[name] <= 16384:
            raise ValueError('Invalid source dimensions')
    clip_range = job.get('range', {})
    start, end = clip_range.get('start'), clip_range.get('end')
    if not all(isinstance(v, (int, float)) and math.isfinite(v) for v in (start, end)):
        raise ValueError('Invalid clip range')
    if start < 0 or not 0 < end - start <= 30:
        raise ValueError('Range must be within 30 seconds')
    rate = job.get('sampleFps', 12)
    if not isinstance(rate, (int, float)) or not 1 <= rate <= 24:
        raise ValueError('Invalid sample FPS')
    points = job.get('points', [])
    ids = [p.get('id') for p in points]
    if not 2 <= len(points) <= 6 or len(set(ids)) != len(ids) or not set(ids) <= ALLOWED_POINTS:
        raise ValueError('Invalid point IDs')
    if not {'bow_a', 'bow_b'} <= set(ids):
        raise ValueError('Both bow axis points are required')
    for group in ('bow', 'bridge', 'string'):
        if (f'{group}_a' in ids) != (f'{group}_b' in ids):
            raise ValueError('Points must be complete pairs')
    for p in points:
        if not all(isinstance(p.get(k), (int, float)) and math.isfinite(p[k]) and 0 <= p[k] <= 1 for k in ('x', 'y')):
            raise ValueError('Invalid normalized seed coordinate')


def run(video_path: Path, job: dict, checkpoint: Path, tapnet_root: Path, device_name: str, output: Path, *, max_frames: int | None = None) -> dict:
    validate_job(job)
    if sha256(video_path) != job['source']['sha256']:
        raise ValueError('Video bytes do not match the browser job')
    if not checkpoint.is_file() or checkpoint.stat().st_size < 100_000_000:
        raise ValueError('Official checkpoint missing or incomplete')
    if not (tapnet_root / 'tapnet/tapnext/tapnext_torch.py').is_file():
        raise ValueError('Pass the official tapnet repository root')
    checkpoint_hash = sha256(checkpoint)
    if checkpoint_hash != CHECKPOINT_SHA256:
        raise ValueError('Checkpoint SHA-256 mismatch: use the documented official 256px checkpoint')
    sys.path.insert(0, str(tapnet_root.resolve()))
    import cv2
    import numpy as np
    import torch
    from tapnet.tapnext.tapnext_torch import TAPNext
    from tapnet.tapnext.tapnext_torch_utils import tracker_certainty
    from tapnet.tapnextpp.votsp2026 import utils

    torch.set_num_threads(min(4, torch.get_num_threads()))
    device = torch.device('cuda' if device_name == 'auto' and torch.cuda.is_available() else 'cpu' if device_name == 'auto' else device_name)
    if device.type == 'cuda' and not torch.cuda.is_available():
        raise ValueError('CUDA is unavailable; choose --device cpu')
    print(f'Loading TAPNext++ on {device}; CPU inference can be slow.', flush=True)
    # mmap avoids loading optimizer buffers from the research checkpoint into RAM.
    checkpoint_data = torch.load(checkpoint, map_location='cpu', weights_only=True, mmap=True)
    state_dict = checkpoint_data.get('state_dict', checkpoint_data)
    weights = {key.removeprefix('tapnext.'): value for key, value in state_dict.items()}
    model = TAPNext(image_size=(256, 256))
    model.load_state_dict(weights)
    del weights, state_dict, checkpoint_data
    model = model.to(device).eval()
    capture = cv2.VideoCapture(str(video_path))
    if not capture.isOpened():
        raise ValueError('Video cannot be decoded')
    width = int(capture.get(cv2.CAP_PROP_FRAME_WIDTH))
    height = int(capture.get(cv2.CAP_PROP_FRAME_HEIGHT))
    source_fps = float(capture.get(cv2.CAP_PROP_FPS))
    source_frames = int(capture.get(cv2.CAP_PROP_FRAME_COUNT))
    if (width, height) != (job['source']['width'], job['source']['height']):
        capture.release()
        raise ValueError('Decoded dimensions differ (including possible rotation); re-export an upright MP4')
    if not math.isfinite(source_fps) or source_fps <= 0:
        capture.release()
        raise ValueError('Video frame timing is unavailable')
    start, end = job['range']['start'], job['range']['end']
    if source_frames and end > source_frames / source_fps + .04:
        capture.release()
        raise ValueError('Requested range is outside the video')
    sample_fps = min(job.get('sampleFps', 12), source_fps)
    # Seed on the exact frame containing the requested start time (floor, not round).
    first_index = max(0, int(math.floor(start * source_fps + 1e-7)))
    indices = sorted(set(first_index + int(math.floor(i * source_fps / sample_fps + 1e-7)) for i in range(math.ceil((end-start)*sample_fps))))
    indices = [i for i in indices if i / source_fps < end]
    if max_frames is not None:
        indices = indices[:max_frames]
    if len(indices) < 2:
        capture.release()
        raise ValueError('At least two source frames are required')
    points = job['points']
    normalized_xy = np.asarray([[p['x'], p['y']] for p in points], dtype=np.float32)
    display_xy = normalized_xy * np.asarray([width, height], dtype=np.float32)
    queries = utils.make_query_tensor(utils.display_to_model(display_xy, height, width), device)
    frames = []
    state = None
    begun = time.perf_counter()
    try:
        with torch.inference_mode():
            for index in indices:
                capture.set(cv2.CAP_PROP_POS_FRAMES, index)
                ok, frame_bgr = capture.read()
                if not ok:
                    raise ValueError(f'Cannot decode source frame {index}')
                timestamp = capture.get(cv2.CAP_PROP_POS_MSEC) / 1000.0
                timing_status = 'MEASURED'
                if not math.isfinite(timestamp) or (frames and timestamp <= frames[-1].get('sourceTimestampSeconds', -1)):
                    timestamp = index/source_fps
                    timing_status = 'ESTIMATED'
                if timestamp >= end:
                    break
                tensor = utils.preprocess_frame(frame_bgr, device, 256)
                with torch.amp.autocast(device_type=device.type, enabled=device.type == 'cuda', dtype=torch.float16 if device.type == 'cuda' else torch.bfloat16):
                    tracks, track_logits, visibility_logits, state = model(video=tensor, query_points=queries if state is None else None, state=state)
                certainty = tracker_certainty(tracks.float(), track_logits.float())
                visibility = torch.sigmoid(visibility_logits.float())
                confidence = (certainty * visibility)[0, 0, :, 0].cpu().numpy()
                # Official TAPNext query / output layout is [t,y,x] / [y,x].
                xy = tracks[0, 0].float().cpu().numpy()[:, ::-1] / 256.0
                row = {'t': max(start, timestamp), 'sourceFrame': index, 'sourceTimestampSeconds': timestamp, 'timingStatus': timing_status, 'points': {}}
                for i, point in enumerate(points):
                    x, y = float(xy[i, 0]), float(xy[i, 1])
                    conf = float(np.clip(confidence[i], 0, 1))
                    valid = math.isfinite(x+y+conf) and 0 <= x <= 1 and 0 <= y <= 1 and conf >= .5
                    row['points'][point['id']] = {'x': x if math.isfinite(x) else None, 'y': y if math.isfinite(y) else None, 'confidence': conf if math.isfinite(conf) else 0, 'valid': valid, 'reason': None if valid else 'OCCLUDED_OR_UNCERTAIN', 'status': 'ESTIMATED' if valid else 'UNKNOWN'}
                frames.append(row)
                print(f'{len(frames)}/{len(indices)} frames, {time.perf_counter()-begun:.1f}s elapsed', flush=True)
    finally:
        capture.release()
    try:
        actual_commit = subprocess.check_output(['git', '-C', str(tapnet_root), 'rev-parse', 'HEAD'], text=True).strip()
    except (OSError, subprocess.CalledProcessError):
        actual_commit = 'unknown'
    result = {'schema': 'bowflow-point-tracks-v1', 'jobId': job['jobId'], 'engine': 'tapnext++', 'source': job['source'], 'range': job['range'], 'frames': frames, 'provenance': {'upstreamCommit': actual_commit, 'expectedUpstreamCommit': UPSTREAM_COMMIT, 'checkpointSha256': checkpoint_hash, 'torchVersion': torch.__version__, 'device': str(device), 'inputResolution': 256, 'sourceFps': source_fps, 'sampleFps': sample_fps, 'sampledFrames': len(frames), 'partialSmokeTest': max_frames is not None, 'elapsedSeconds': time.perf_counter()-begun}}
    output.parent.mkdir(parents=True, exist_ok=True)
    temporary = output.with_suffix(output.suffix + '.tmp')
    temporary.write_text(json.dumps(result, ensure_ascii=False, allow_nan=False), encoding='utf-8')
    temporary.replace(output)
    return result


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--video', type=Path, required=True)
    parser.add_argument('--job', type=Path, required=True)
    parser.add_argument('--checkpoint', type=Path, required=True)
    parser.add_argument('--tapnet-root', type=Path, required=True)
    parser.add_argument('--output', type=Path, default=Path('bowflow-tapnext-tracks.json'))
    parser.add_argument('--device', choices=['auto', 'cpu', 'cuda'], default='auto')
    parser.add_argument('--max-frames', type=int, help='Engineering smoke test only; results are marked partial')
    args = parser.parse_args()
    try:
        result = run(args.video, json.loads(args.job.read_text(encoding='utf-8')), args.checkpoint, args.tapnet_root, args.device, args.output, max_frames=args.max_frames)
    except (ValueError, OSError, RuntimeError) as error:
        raise SystemExit(str(error)) from error
    print(f"Saved {len(result['frames'])} real TAPNext++ predictions to {args.output}")


if __name__ == '__main__':
    main()
