from __future__ import annotations

from dataclasses import dataclass
from typing import Dict, List, Tuple


PAD = "<PAD>"
UNK = "<UNK>"


@dataclass
class Vocab:
    stoi: Dict[str, int]
    itos: Dict[int, str]

    @classmethod
    def build(cls, corpus_tokens: List[List[str]], min_freq: int = 1) -> "Vocab":
        freq: Dict[str, int] = {}
        for sent in corpus_tokens:
            for w in sent:
                freq[w] = freq.get(w, 0) + 1
        stoi: Dict[str, int] = {PAD: 0, UNK: 1}
        for w, c in sorted(freq.items(), key=lambda kv: (-kv[1], kv[0])):
            if c >= min_freq and w not in stoi:
                stoi[w] = len(stoi)
        itos = {i: w for w, i in stoi.items()}
        return cls(stoi=stoi, itos=itos)

    def encode(self, tokens: List[str]) -> List[int]:
        return [self.stoi.get(w, self.stoi[UNK]) for w in tokens]

    def decode(self, ids: List[int]) -> List[str]:
        return [self.itos.get(i, UNK) for i in ids]


def tokenize(sentences: List[str]) -> List[List[str]]:
    return [s.strip().split() for s in sentences]


def pad_sequences(batch: List[List[int]], pad_id: int, max_len: int | None = None) -> List[List[int]]:
    if max_len is None:
        max_len = max((len(x) for x in batch), default=0)
    out: List[List[int]] = []
    for seq in batch:
        seq = list(seq)
        if len(seq) < max_len:
            seq = seq + [pad_id] * (max_len - len(seq))
        else:
            seq = seq[:max_len]
        out.append(seq)
    return out


def preprocess_parallel(en_sentences: List[str], fr_sentences: List[str], ref_length: int | None = None) -> Tuple[List[List[int]], List[List[int]], Vocab, Vocab]:
    en_tok = tokenize(en_sentences)
    fr_tok = tokenize(fr_sentences)
    en_vocab = Vocab.build(en_tok)
    fr_vocab = Vocab.build(fr_tok)
    en_ids = [en_vocab.encode(x) for x in en_tok]
    fr_ids = [fr_vocab.encode(x) for x in fr_tok]
    en_pad = pad_sequences(en_ids, pad_id=en_vocab.stoi[PAD])
    fr_pad = pad_sequences(fr_ids, pad_id=fr_vocab.stoi[PAD], max_len=ref_length)
    if ref_length is not None:
        en_pad = pad_sequences(en_pad, pad_id=en_vocab.stoi[PAD], max_len=ref_length)
    return en_pad, fr_pad, en_vocab, fr_vocab

