"""第 20—23 课。模型提供字的分数；程序只约束外形并公开候选评分。"""
import math
import re
from pathlib import Path

import torch

from .common import DEFAULT_DATA, device_for, digest, read_jsonl
from .data import HAN, THEMES, normalize
from .train import load_model, seed_all


def parse_keywords(value):
    words = [normalize(x) for x in re.split(r"[,，\s]+", value) if x.strip()]
    words = list(dict.fromkeys(words))
    if len(words) > 3: raise ValueError("一次最多输入三个关键词，便于小模型集中表达。")
    if any(not HAN.fullmatch(x) or len(x)>4 for x in words):
        raise ValueError("关键词请使用 1—4 个汉字，例如：明月、秋雨、故乡。")
    return words


def keyword_score(text, words):
    if not words: return 1.0
    values = []
    for word in words:
        if word in text: values.append(1.0)
        elif word in THEMES and any(x in text for x in THEMES[word]): values.append(0.4)
        else: values.append(0.3*sum(x in text for x in word)/len(word))
    return sum(values)/len(values)


class Writer:
    def __init__(self, checkpoint, device="auto", data_dir=DEFAULT_DATA):
        self.device = device_for(device)
        if self.device.type == "cpu": torch.set_num_threads(4)
        self.model, self.tokenizer, self.metadata = load_model(checkpoint, self.device)
        self.char_ids = torch.tensor([i for i,t in enumerate(self.tokenizer.tokens) if len(t)==1 and HAN.fullmatch(t)], device=self.device)
        self.known = set()
        self.corpus_check_note = "未找到完整且版本匹配的配套材料，未执行原诗重复检查。"
        paths = {split: Path(data_dir) / f"{split}.jsonl" for split in ("train", "val", "test")}
        if all(path.is_file() for path in paths.values()):
            actual = {split: digest(path.read_text(encoding="utf-8")) for split,path in paths.items()}
            if actual == self.metadata.get("data_fingerprint"):
                for path in paths.values(): self.known.update(r["id"] for r in read_jsonl(path))
                self.corpus_check_note = "已对与模型版本匹配的三份材料执行整首精确重复检查。"

    @torch.no_grad()
    def write(self, start="", keywords=(), form=5, count=1, temperature=0.8, top_k=40, candidates=12, seed=2026):
        if form not in (5, 7): raise ValueError("目前支持五言和七言。")
        if count < 1 or count > 20: raise ValueError("每次生成数量应在 1—20 之间。")
        if candidates < 1 or candidates > 64: raise ValueError("候选数量应在 1—64 之间。")
        if temperature <= 0 or not math.isfinite(temperature): raise ValueError("temperature 必须是有限正数。")
        if top_k < 1: raise ValueError("top-k 必须大于 0。")
        start = normalize(start)
        if start and (not HAN.fullmatch(start) or len(start)>form): raise ValueError(f"开头请使用不超过 {form} 个汉字。")
        self.tokenizer.encode(start)
        for word in keywords: self.tokenizer.encode(word)
        prefix = self.tokenizer.prefix(form, keywords)
        if len(prefix)+4*(form+1)>self.model.config.context:
            raise ValueError("输入过长，请缩短关键词。")
        seed_all(seed)
        pool = []
        generated = set()
        for attempt in range(max(2, math.ceil(count / candidates)+1)):
            sequences = torch.tensor([prefix+self.tokenizer.encode(start)]*candidates, device=self.device)
            logprobs = torch.zeros(candidates, device=self.device)
            for position in range(len(start), 4*(form+1)):
                if position % (form+1) == form:
                    punctuation = "，" if position//(form+1) % 2 == 0 else "。"
                    chosen = torch.full((candidates,1), self.tokenizer.ids[punctuation], device=self.device, dtype=torch.long)
                else:
                    logits, _ = self.model(sequences)
                    scores = logits[:, -1, self.char_ids]
                    values, indices = torch.topk(scores / temperature, min(top_k,len(self.char_ids)), dim=-1)
                    picked = torch.multinomial(values.softmax(dim=-1), 1)
                    local = indices.gather(1,picked)
                    chosen = self.char_ids[local]
                    logprobs += scores.log_softmax(dim=-1).gather(1,local).squeeze(1)
                sequences = torch.cat([sequences, chosen], dim=1)
            scores_cpu = logprobs.cpu().tolist()
            for i, sequence in enumerate(sequences[:,len(prefix):].cpu().tolist()):
                text = self.tokenizer.decode(sequence)
                plain = text.replace("，", "").replace("。", "")
                duplicate = digest(plain) in self.known
                if text in generated or duplicate: continue
                generated.add(text)
                repeated = sum(max(0,plain.count(ch)-2) for ch in set(plain))
                lines = re.findall(r"[^，。]+[，。]",text)
                bigrams = [line[j:j+2] for line in lines for j in range(len(line)-2)]
                repeated_phrases = len(bigrams)-len(set(bigrams))
                repeated_endings = len(lines)-len({line[-2] for line in lines})
                doubled = sum(a==b for line in lines for a,b in zip(line[:-2],line[1:-1]))
                theme = keyword_score(text, keywords)
                fluency = scores_cpu[i] / max(1, 4*form-len(start))
                score = fluency + 2.0*theme - 0.35*repeated - 0.45*repeated_phrases - 0.7*repeated_endings - 0.7*doubled
                pool.append({"text": text, "lines": lines, "form": form,
                             "model_step": self.metadata["step"],
                             "keywords": list(keywords), "keyword_literal_hits": [x for x in keywords if x in text],
                             "keyword_rule_score": round(theme,3), "model_mean_log_probability": round(fluency,3),
                             "selection_score": round(score,3), "corpus_exact_match": duplicate if self.known else None,
                             "repeated_bigrams": repeated_phrases, "repeated_line_endings": repeated_endings,
                             "corpus_check_available": bool(self.known), "corpus_check_note": self.corpus_check_note})
            if len(pool) >= count: break
        if len(pool)<count:
            raise ValueError("候选诗重复过多，未凑齐所需数量。请增加 temperature 或 candidates。")
        pool.sort(key=lambda x:x["selection_score"], reverse=True)
        return pool[:count]
