from __future__ import annotations

import re
import unicodedata
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any


@dataclass(frozen=True)
class ContextBudgetMetrics:
    input_count: int
    output_count: int
    duplicates_removed: int = 0
    truncated_items: int = 0
    truncated_chars: int = 0


def normalize_budget_text(value: Any) -> str:
    return re.sub(r"\s+", " ", str(value or "")).strip()


def normalize_budget_key(value: Any) -> str:
    normalized = unicodedata.normalize("NFKD", normalize_budget_text(value))
    return "".join(char for char in normalized if not unicodedata.combining(char)).casefold()


def deduplicate_texts(values: list[str]) -> tuple[list[str], ContextBudgetMetrics]:
    unique: list[str] = []
    seen: set[str] = set()
    for value in values or []:
        text = normalize_budget_text(value)
        if not text:
            continue
        key = normalize_budget_key(text)
        if key in seen:
            continue
        seen.add(key)
        unique.append(text)
    metrics = ContextBudgetMetrics(
        input_count=len(values or []),
        output_count=len(unique),
        duplicates_removed=max(0, len(values or []) - len(unique)),
    )
    return unique, metrics


def deduplicate_records[T](
    values: list[T],
    *,
    key_fn: Callable[[T], Any],
) -> tuple[list[T], ContextBudgetMetrics]:
    unique: list[T] = []
    seen: set[str] = set()
    for value in values or []:
        key = normalize_budget_key(key_fn(value))
        if not key:
            continue
        if key in seen:
            continue
        seen.add(key)
        unique.append(value)
    metrics = ContextBudgetMetrics(
        input_count=len(values or []),
        output_count=len(unique),
        duplicates_removed=max(0, len(values or []) - len(unique)),
    )
    return unique, metrics


def truncate_text(value: str, *, max_chars: int | None) -> tuple[str, ContextBudgetMetrics]:
    text = str(value or "")
    if not max_chars or max_chars <= 0 or len(text) <= max_chars:
        return text, ContextBudgetMetrics(input_count=1, output_count=1)
    truncated = text[:max_chars].rstrip()
    return truncated, ContextBudgetMetrics(
        input_count=1,
        output_count=1,
        truncated_items=1,
        truncated_chars=len(text) - len(truncated),
    )


def limit_items[T](values: list[T], *, max_items: int | None) -> tuple[list[T], ContextBudgetMetrics]:
    if not max_items or max_items <= 0 or len(values or []) <= max_items:
        return list(values or []), ContextBudgetMetrics(
            input_count=len(values or []),
            output_count=len(values or []),
        )
    limited = list(values[:max_items])
    return limited, ContextBudgetMetrics(
        input_count=len(values),
        output_count=len(limited),
    )
