"""Citation support scoring: token overlap + number agreement, no model.

`support(claim, chunk)` and `check(answer_with_cites, chunks)` below are what you implement.
Everything else here is plumbing: tokenising, the stopword list, and the `[n]` citation parser.
"""

from __future__ import annotations

import re
from collections.abc import Iterator
from dataclasses import dataclass, field

_TOKEN_RE = re.compile(r"[a-z0-9]+")
_NUMBER_RE = re.compile(r"\d+(?:\.\d+)?%?")

STOPWORDS = frozenset(
    {
        "the", "a", "an", "and", "or", "of", "to", "in", "is", "are", "was", "were",
        "for", "on", "with", "as", "by", "at", "from", "that", "this", "it", "its",
        "be", "been", "has", "have", "had", "not", "no", "but", "which", "into",
        "than", "then", "so", "such", "also", "if", "can", "could", "will", "would",
        "across", "using", "approximately",
    }
)

NUMBER_MISMATCH_PENALTY = 0.3
SUPPORT_THRESHOLD = 0.5


def _tokens(text: str) -> list[str]:
    return _TOKEN_RE.findall(text.lower())


def _content_tokens(text: str) -> set[str]:
    """Lowercase word tokens with STOPWORDS removed -- what `support` overlaps on."""
    return {t for t in _tokens(text) if t not in STOPWORDS}


def _numbers(text: str) -> set[str]:
    """Every number written in `text`, as the exact substring matched (so "1889" != "1889.0")."""
    return set(_NUMBER_RE.findall(text))


@dataclass
class ClaimResult:
    claim: str
    cited_chunk: int
    status: str
    score: float


@dataclass
class Report:
    results: list[ClaimResult] = field(default_factory=list)

    @property
    def statuses(self) -> list[str]:
        return [r.status for r in self.results]


_CITE_RE = re.compile(r"^(.*?)\s*\[(\d+)\]\s*$")


def _parse_claims(answer_with_cites: str) -> Iterator[tuple[str, int]]:
    """Yield (claim_text, cited_index) for each non-blank line of `answer_with_cites`; the index
    is 1-indexed, exactly as it reads in the `[n]` marker -- `check` is the one that turns it
    into a list index."""
    for raw_line in answer_with_cites.strip().splitlines():
        line = raw_line.strip()
        if not line:
            continue
        m = _CITE_RE.match(line)
        if not m:
            raise ValueError(f"claim has no trailing citation marker like [1]: {line!r}")
        yield m.group(1).strip(), int(m.group(2))


def support(claim: str, chunk: str) -> float:
    """A score in [0, 1]: the fraction of `claim`'s content tokens (see `_content_tokens`) that
    also appear in `chunk`'s content tokens. 0.0 when the claim has no content tokens at all --
    never divide by zero. Then, if the claim states a number (`_numbers`) that is not a subset of
    the chunk's numbers, multiply the score by NUMBER_MISMATCH_PENALTY: a claim can share every
    word with its chunk and still be citing a number the chunk never says."""
    ...


def check(answer_with_cites: str, chunks: list[str]) -> Report:
    """Parse `answer_with_cites` with `_parse_claims`, and for each (claim, 1-indexed n):

    - n must satisfy 1 <= n <= len(chunks); otherwise raise ValueError.
    - chunk = chunks[n - 1] -- the marker is 1-indexed, the list is not.
    - status is "unsupported" when the *word* overlap alone (ignore the number penalty here) is
      below SUPPORT_THRESHOLD; otherwise "supported" when the claim's numbers are all backed by
      the chunk, "wrong-number" when they are not.
    - `score` on the ClaimResult is `support(claim, chunk)` (the penalized score).

    Return a Report holding one ClaimResult per claim, in order.
    """
    ...
