Operations

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

src/rag_evolution/operations.py · 302 lines · sha256 718505d2facf…

"""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)