"""Fill non-English translation packages with checkpointed NLLB machine drafts.

Requires the project-local .translation-runtime and downloads the configured model
to .translation-models on first use. Drafts are never marked native-approved.
"""

from __future__ import annotations

import json
import csv
import re
import sys
from pathlib import Path

REPO = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(REPO / ".translation-runtime"))

import torch
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer

ROOT = REPO / "content" / "speech"
PLAN = json.loads((ROOT / "plan" / "content-plan.json").read_text(encoding="utf-8"))
MODEL_ID = "facebook/nllb-200-distilled-600M"
MODEL_CACHE = REPO / ".translation-models"
MEMORY_PATH = ROOT / "tracking" / "machine-translation-memory.json"
LANG = {"hi":"hin_Deva","ta":"tam_Taml","te":"tel_Telu","kn":"kan_Knda","ml":"mal_Mlym","mr":"mar_Deva","bn":"ben_Beng"}
SLOT = re.compile(r"\{[^{}]+\}")


def concept_path(c): return ROOT / "concepts" / Path(*c["category"].split(".")) / c["key"]
def protect(text):
    slots=[]
    def repl(m): slots.append(m.group(0)); return f"ZXQSLOT{len(slots)-1}QXZ"
    return SLOT.sub(repl,text),slots
def restore(text,slots):
    for i,value in enumerate(slots):
        token=f"ZXQSLOT{i}QXZ"
        text=text.replace(token,value).replace(token.lower(),value)
    return text.strip()

memory=json.loads(MEMORY_PATH.read_text(encoding="utf-8")) if MEMORY_PATH.exists() else {"schema_version":1,"model":MODEL_ID,"translations":{}}
packages=[]; needed={code:set() for code in LANG}
for c in PLAN["concepts"]:
    for code in LANG:
        p=concept_path(c)/"locales"/code/"translation-package.json"; d=json.loads(p.read_text(encoding="utf-8")); packages.append((p,d,code))
        texts=[d["english_reference"]["preferred_label"],*d["english_reference"]["natural_alternates"]]
        texts += [x["source_meaning_en"] for x in d["target_content"]["expressions"] if x["source_meaning_en"]]
        texts += [x["source_meaning_en"] for x in d["target_content"]["caregiver_directions"] if x["source_meaning_en"]]
        needed[code].update(texts)

tokenizer=AutoTokenizer.from_pretrained(MODEL_ID,cache_dir=MODEL_CACHE,src_lang="eng_Latn")
model=AutoModelForSeq2SeqLM.from_pretrained(MODEL_ID,cache_dir=MODEL_CACHE)
model.eval(); torch.set_num_threads(max(1,min(8,torch.get_num_threads())))

for code,nllb in LANG.items():
    lang_mem=memory["translations"].setdefault(code,{})
    todo=sorted(t for t in needed[code] if t not in lang_mem)
    print(f"{code}: {len(todo)} new strings",flush=True)
    for start in range(0,len(todo),32):
        source=todo[start:start+32]; protected=[]; slots=[]
        for text in source:
            p,s=protect(text); protected.append(p); slots.append(s)
        encoded=tokenizer(protected,return_tensors="pt",padding=True,truncation=True,max_length=256)
        with torch.inference_mode():
            output=model.generate(**encoded,forced_bos_token_id=tokenizer.convert_tokens_to_ids(nllb),max_length=256,num_beams=1)
        translated=tokenizer.batch_decode(output,skip_special_tokens=True)
        for source_text,target_text,slot_values in zip(source,translated,slots): lang_mem[source_text]=restore(target_text,slot_values)
        MEMORY_PATH.write_text(json.dumps(memory,ensure_ascii=False,indent=2)+"\n",encoding="utf-8")
        print(f"{code}: {min(start+32,len(todo))}/{len(todo)}",flush=True)

for path,d,code in packages:
    mem=memory["translations"][code]; target=d["target_content"]
    target["preferred_term"]=mem[d["english_reference"]["preferred_label"]]
    target["aac_short_label"]=target["preferred_term"]
    target["natural_alternatives"]=[mem[x] for x in d["english_reference"]["natural_alternates"]]
    for x in target["expressions"]:
        if x["source_meaning_en"]:
            x["target_text"]=mem[x["source_meaning_en"]]; x["status"]="machine_translation_draft_needs_native_review"
    for x in target["caregiver_directions"]:
        if x["source_meaning_en"]:
            x["target_text"]=mem[x["source_meaning_en"]]; x["status"]="machine_translation_draft_needs_native_review"
    d["machine_translation"]={"model":MODEL_ID,"source_language":"eng_Latn","target_language":LANG[code],"status":"draft_only","native_review_required":True}
    d["review"]["status"]="machine_translation_draft_needs_native_review"
    path.write_text(json.dumps(d,ensure_ascii=False,indent=2)+"\n",encoding="utf-8")

queue_path = ROOT / "tracking" / "translation-queue.csv"
with queue_path.open("r", encoding="utf-8", newline="") as stream:
    rows = list(csv.DictReader(stream))
fieldnames = list(rows[0].keys()) if rows else []
for row in rows:
    if row["language"] in LANG:
        row["translation_status"] = "machine_translation_draft_needs_native_review"
with queue_path.open("w", encoding="utf-8", newline="") as stream:
    writer = csv.DictWriter(stream, fieldnames=fieldnames)
    writer.writeheader()
    writer.writerows(rows)

print(f"Updated {len(packages)} translation packages",flush=True)
