"""Production measurements, release fingerprints, budgets, and Pareto analysis."""

import hashlib
import json
import math
from dataclasses import asdict, dataclass
from typing import Any, Dict, Iterable, List, Mapping, Sequence, Tuple


@dataclass(frozen=True)
class StageMeasurement:
    """One request-stage observation in common operational units."""

    stage: str
    duration_ms: float
    cost_usd: float = 0.0
    input_tokens: int = 0
    output_tokens: int = 0
    calls: int = 1
    cache_hit: bool = False
    failed: bool = False

    def __post_init__(self) -> None:
        if not self.stage.strip():
            raise ValueError("stage must be non-empty")
        if self.duration_ms < 0 or self.cost_usd < 0:
            raise ValueError("duration and cost must be non-negative")
        if min(self.input_tokens, self.output_tokens, self.calls) < 0:
            raise ValueError("tokens and calls must be non-negative")


@dataclass(frozen=True)
class RequestMeasurement:
    """End-to-end trace with quality and safety outcomes."""

    request_id: str
    stages: Tuple[StageMeasurement, ...]
    quality: float
    cited: bool
    abstained: bool
    safe: bool = True

    @property
    def latency_ms(self) -> float:
        return sum(stage.duration_ms for stage in self.stages)

    @property
    def cost_usd(self) -> float:
        return sum(stage.cost_usd for stage in self.stages)

    @property
    def failed(self) -> bool:
        return any(stage.failed for stage in self.stages)


@dataclass(frozen=True)
class ServiceBudget:
    """Hard per-request constraints independent of a learned search policy."""

    maximum_latency_ms: float
    maximum_cost_usd: float
    maximum_retrieval_calls: int
    maximum_generation_tokens: int

    def __post_init__(self) -> None:
        if min(
            self.maximum_latency_ms,
            self.maximum_cost_usd,
            self.maximum_retrieval_calls,
            self.maximum_generation_tokens,
        ) < 0:
            raise ValueError("budgets must be non-negative")


@dataclass(frozen=True)
class BudgetCheck:
    """Budget decision and every violated constraint."""

    allowed: bool
    violations: Tuple[str, ...]
    latency_ms: float
    cost_usd: float
    retrieval_calls: int
    generation_tokens: int


@dataclass(frozen=True)
class SystemCandidate:
    """One evaluated configuration for constrained/Pareto selection."""

    name: str
    quality: float
    cost_usd: float
    latency_ms: float
    risk: float

    def __post_init__(self) -> None:
        if not self.name:
            raise ValueError("candidate name must be non-empty")
        if min(self.cost_usd, self.latency_ms, self.risk) < 0:
            raise ValueError("cost, latency, and risk must be non-negative")


def percentile(values: Sequence[float], probability: float) -> float:
    """Linearly interpolated percentile using a deterministic rank convention."""

    if not values:
        raise ValueError("percentile requires observations")
    if not 0.0 <= probability <= 1.0:
        raise ValueError("probability must be between zero and one")
    ordered = sorted(values)
    position = probability * (len(ordered) - 1)
    lower = int(math.floor(position))
    upper = int(math.ceil(position))
    if lower == upper:
        return ordered[lower]
    fraction = position - lower
    return ordered[lower] * (1.0 - fraction) + ordered[upper] * fraction


def latency_summary(requests: Sequence[RequestMeasurement]) -> Mapping[str, float]:
    """End-to-end p50/p95/p99 plus success and cache rates."""

    if not requests:
        return {}
    latencies = [request.latency_ms for request in requests]
    stages = [stage for request in requests for stage in request.stages]
    cacheable = [stage for stage in stages if stage.stage in {"retrieve", "rerank", "generate"}]
    return {
        "p50_latency_ms": percentile(latencies, 0.50),
        "p95_latency_ms": percentile(latencies, 0.95),
        "p99_latency_ms": percentile(latencies, 0.99),
        "mean_cost_usd": sum(request.cost_usd for request in requests) / len(requests),
        "failure_rate": sum(request.failed for request in requests) / len(requests),
        "safe_rate": sum(request.safe for request in requests) / len(requests),
        "citation_rate": sum(request.cited for request in requests) / len(requests),
        "abstention_rate": sum(request.abstained for request in requests) / len(requests),
        "cache_hit_rate": (
            sum(stage.cache_hit for stage in cacheable) / len(cacheable) if cacheable else 0.0
        ),
    }


def stage_summary(requests: Sequence[RequestMeasurement]) -> Mapping[str, Mapping[str, float]]:
    """Attribute latency, cost, calls, and errors to individual stages."""

    grouped: Dict[str, List[StageMeasurement]] = {}
    for request in requests:
        for stage in request.stages:
            grouped.setdefault(stage.stage, []).append(stage)
    output: Dict[str, Mapping[str, float]] = {}
    for name, observations in sorted(grouped.items()):
        durations = [item.duration_ms for item in observations]
        output[name] = {
            "count": float(len(observations)),
            "p50_ms": percentile(durations, 0.50),
            "p95_ms": percentile(durations, 0.95),
            "mean_cost_usd": sum(item.cost_usd for item in observations) / len(observations),
            "mean_calls": sum(item.calls for item in observations) / len(observations),
            "error_rate": sum(item.failed for item in observations) / len(observations),
        }
    return output


def check_budget(request: RequestMeasurement, budget: ServiceBudget) -> BudgetCheck:
    """Check hard limits after a request or simulated trajectory."""

    retrieval_calls = sum(
        stage.calls for stage in request.stages if stage.stage in {"retrieve", "search", "tool"}
    )
    generation_tokens = sum(
        stage.output_tokens for stage in request.stages if stage.stage in {"generate", "answer"}
    )
    violations = []
    if request.latency_ms > budget.maximum_latency_ms:
        violations.append("latency")
    if request.cost_usd > budget.maximum_cost_usd:
        violations.append("cost")
    if retrieval_calls > budget.maximum_retrieval_calls:
        violations.append("retrieval_calls")
    if generation_tokens > budget.maximum_generation_tokens:
        violations.append("generation_tokens")
    return BudgetCheck(
        allowed=not violations,
        violations=tuple(violations),
        latency_ms=request.latency_ms,
        cost_usd=request.cost_usd,
        retrieval_calls=retrieval_calls,
        generation_tokens=generation_tokens,
    )


def utility(
    candidate: SystemCandidate,
    cost_weight: float = 1.0,
    latency_weight: float = 0.001,
    risk_weight: float = 1.0,
) -> float:
    """Constrained-system objective ``quality - λc cost - λl latency - λr risk``."""

    if min(cost_weight, latency_weight, risk_weight) < 0:
        raise ValueError("utility weights must be non-negative")
    return (
        candidate.quality
        - cost_weight * candidate.cost_usd
        - latency_weight * candidate.latency_ms
        - risk_weight * candidate.risk
    )


def dominates(left: SystemCandidate, right: SystemCandidate) -> bool:
    """Whether ``left`` is no worse on every objective and better on one."""

    no_worse = (
        left.quality >= right.quality
        and left.cost_usd <= right.cost_usd
        and left.latency_ms <= right.latency_ms
        and left.risk <= right.risk
    )
    strictly_better = (
        left.quality > right.quality
        or left.cost_usd < right.cost_usd
        or left.latency_ms < right.latency_ms
        or left.risk < right.risk
    )
    return no_worse and strictly_better


def pareto_frontier(candidates: Sequence[SystemCandidate]) -> Tuple[SystemCandidate, ...]:
    """Return non-dominated configurations in stable name order."""

    frontier = [
        candidate
        for candidate in candidates
        if not any(
            dominates(other, candidate)
            for other in candidates
            if other.name != candidate.name
        )
    ]
    return tuple(sorted(frontier, key=lambda item: item.name))


def constrained_choice(
    candidates: Sequence[SystemCandidate],
    minimum_quality: float,
    maximum_cost_usd: float,
    maximum_latency_ms: float,
    maximum_risk: float,
) -> Tuple[SystemCandidate, ...]:
    """Filter candidates through release gates before optimizing a mean score."""

    feasible = [
        candidate
        for candidate in candidates
        if candidate.quality >= minimum_quality
        and candidate.cost_usd <= maximum_cost_usd
        and candidate.latency_ms <= maximum_latency_ms
        and candidate.risk <= maximum_risk
    ]
    feasible.sort(key=lambda item: (-item.quality, item.cost_usd, item.latency_ms, item.name))
    return tuple(feasible)


def configuration_fingerprint(configuration: Mapping[str, Any]) -> str:
    """Content-address a complete, JSON-serializable experiment configuration."""

    serialized = json.dumps(configuration, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
    return hashlib.sha256(serialized.encode("utf-8")).hexdigest()


def release_manifest(
    corpus_snapshot: str,
    parser_version: str,
    chunker_version: str,
    embedding_version: str,
    index_version: str,
    reranker_version: str,
    generator_version: str,
    prompt_version: str,
) -> Mapping[str, str]:
    """Create a minimal immutable release identity and its aggregate hash."""

    components = {
        "corpus_snapshot": corpus_snapshot,
        "parser_version": parser_version,
        "chunker_version": chunker_version,
        "embedding_version": embedding_version,
        "index_version": index_version,
        "reranker_version": reranker_version,
        "generator_version": generator_version,
        "prompt_version": prompt_version,
    }
    if not all(components.values()):
        raise ValueError("release component versions must be non-empty")
    return {**components, "release_id": configuration_fingerprint(components)}


def request_as_dict(request: RequestMeasurement) -> Mapping[str, Any]:
    """Serialize a trace without losing stage-level measurements."""

    return asdict(request)
