"""Build deterministic train/dev/test subsets without cutting off SQL targets."""
import argparse
import hashlib
import json
import random
from pathlib import Path
from transformers import AutoTokenizer
from sql_task import execute, gold_sql, messages


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--source", type=Path, required=True)
    p.add_argument("--model", required=True)
    p.add_argument("--out", type=Path, required=True)
    p.add_argument("--train-size", type=int, default=2000)
    p.add_argument("--valid-size", type=int, default=100)
    p.add_argument("--test-size", type=int, default=200)
    p.add_argument("--max-length", type=int, default=512)
    p.add_argument("--seed", type=int, default=20260929)
    args = p.parse_args()
    args.out.mkdir(parents=True, exist_ok=True)
    tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=False)
    report = {"seed": args.seed, "max_length": args.max_length, "splits": {}}
    table_ids = {}
    for split, target, count in [("train", "train", args.train_size), ("dev", "valid", args.valid_size), ("test", "test", args.test_size)]:
        tables = {t["id"]: t for t in map(json.loads, (args.source / f"{split}.tables.jsonl").read_text().splitlines())}
        table_ids[split] = set(tables)
        candidates = list(enumerate(map(json.loads, (args.source / f"{split}.jsonl").read_text().splitlines())))
        random.Random(args.seed).shuffle(candidates)
        records, skipped = [], {"length": 0, "gold_error": 0}
        for source_row, q in candidates:
            table = tables[q["table_id"]]
            try:
                sql = gold_sql(q["sql"], table)
                chat = messages(q["question"], table)
                train_chat = chat + [{"role": "assistant", "content": sql}]
                tokens = tokenizer.apply_chat_template(train_chat, return_dict=False)
                prompt = tokenizer.apply_chat_template(chat, add_generation_prompt=True, return_dict=False)
                assert tokens[:len(prompt)] == prompt and len(tokens) > len(prompt)
                if len(tokens) > args.max_length:
                    skipped["length"] += 1
                    continue
                answer = execute(args.source / f"{split}.db", q["table_id"], sql, len(table["types"]))
            except Exception as exc:
                skipped["gold_error"] += 1
                if skipped["gold_error"] <= 3:
                    print("Gold error", split, source_row, str(exc), flush=True)
                if skipped["gold_error"] > 20:
                    raise RuntimeError("Repeated gold execution failure") from exc
                continue
            records.append({"source_split": split, "source_row": source_row, "table_id": q["table_id"], "question": q["question"], "headers": table["header"], "types": table["types"], "query": q["sql"], "messages": chat, "gold_sql": sql, "gold_answer": answer, "tokens": len(tokens), "prompt_tokens": len(prompt)})
            if len(records) % 500 == 0:
                print(split, len(records), "prepared", flush=True)
            if len(records) == count:
                break
        assert len(records) == count
        training = [{"messages": r["messages"] + [{"role": "assistant", "content": r["gold_sql"]}]} for r in records]
        for name, rows in [(f"{target}.jsonl", training), (f"{target}.records.jsonl", records)]:
            (args.out / name).write_text("".join(json.dumps(r, ensure_ascii=False) + "\n" for r in rows))
        lengths = sorted(r["tokens"] for r in records)
        report["splits"][target] = {"source": split, "count": len(records), "unique_tables": len({r['table_id'] for r in records}), "skipped": skipped, "tokens_min": lengths[0], "tokens_median": lengths[len(lengths)//2], "tokens_max": lengths[-1], "sha256": hashlib.sha256((args.out / f"{target}.jsonl").read_bytes()).hexdigest()}
    for a, b in [("train", "dev"), ("train", "test"), ("dev", "test")]:
        assert not table_ids[a] & table_ids[b], (a, b, "overlapping table IDs")
    report["source_table_ids_disjoint"] = True
    (args.out / "manifest.json").write_text(json.dumps(report, indent=2))
    print(json.dumps(report, indent=2))


if __name__ == "__main__":
    main()
