"""Bitemporal evidence, permission-aware lookup, and freshness diagnostics.

RAG systems need at least two clocks:

* **valid time** — when a statement was true in the represented world; and
* **transaction time** — when the retrieval system learned or indexed it.

Conflating these clocks creates hindsight leakage in evaluation and stale or
anachronistic answers in production.  This module provides a small immutable
fact model and an append-only bitemporal store suitable for notebook-scale
experiments.
"""

import math
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Dict, Iterable, List, Optional, Sequence, Set, Tuple


def parse_instant(value: str) -> datetime:
    """Parse an ISO-8601 timestamp and require an explicit timezone."""

    normalized = value[:-1] + "+00:00" if value.endswith("Z") else value
    parsed = datetime.fromisoformat(normalized)
    if parsed.tzinfo is None:
        raise ValueError("timestamps must include an explicit timezone")
    return parsed.astimezone(timezone.utc)


@dataclass(frozen=True)
class TemporalFact:
    """One append-only version of a source claim."""

    source_id: str
    version_id: str
    key: str
    value: str
    valid_from: str
    observed_at: str
    valid_to: str = ""
    retracted_at: str = ""
    principals: Tuple[str, ...] = ()
    trust_domain: str = "unverified"
    provenance: str = ""

    def __post_init__(self) -> None:
        if not all((self.source_id.strip(), self.version_id.strip(), self.key.strip())):
            raise ValueError("source_id, version_id, and key must be non-empty")
        start = parse_instant(self.valid_from)
        observed = parse_instant(self.observed_at)
        if self.valid_to and parse_instant(self.valid_to) <= start:
            raise ValueError("valid_to must be after valid_from")
        if self.retracted_at and parse_instant(self.retracted_at) < observed:
            raise ValueError("retracted_at cannot precede observed_at")

    def valid_at(self, instant: datetime) -> bool:
        """Whether the claim describes the world at ``instant``."""

        start = parse_instant(self.valid_from)
        end = parse_instant(self.valid_to) if self.valid_to else None
        return start <= instant and (end is None or instant < end)

    def known_at(self, instant: datetime) -> bool:
        """Whether the store knew the version at ``instant``."""

        observed = parse_instant(self.observed_at)
        retracted = parse_instant(self.retracted_at) if self.retracted_at else None
        return observed <= instant and (retracted is None or instant < retracted)

    def accessible_to(self, principals: Iterable[str]) -> bool:
        """Public facts have no ACL; restricted facts require an intersection."""

        allowed = set(self.principals)
        return not allowed or bool(allowed & set(principals))


@dataclass(frozen=True)
class TemporalLookup:
    """Lookup result including ambiguity instead of silently choosing a value."""

    key: str
    valid_at: str
    known_at: str
    facts: Tuple[TemporalFact, ...]
    ambiguous: bool
    values: Tuple[str, ...]


class BitemporalStore:
    """Append-only, permission-aware bitemporal fact store."""

    def __init__(self, facts: Sequence[TemporalFact] = ()) -> None:
        self._facts: List[TemporalFact] = []
        self._version_ids: Set[str] = set()
        for fact in facts:
            self.append(fact)

    @property
    def facts(self) -> Tuple[TemporalFact, ...]:
        return tuple(self._facts)

    def append(self, fact: TemporalFact) -> None:
        """Append a unique source version; history is never overwritten."""

        if fact.version_id in self._version_ids:
            raise ValueError("duplicate temporal version: " + fact.version_id)
        self._facts.append(fact)
        self._version_ids.add(fact.version_id)

    def lookup(
        self,
        key: str,
        valid_at: str,
        known_at: str,
        principals: Iterable[str] = (),
        trust_domains: Iterable[str] = (),
    ) -> TemporalLookup:
        """Return every fact simultaneously valid, known, authorized, and trusted."""

        valid_instant = parse_instant(valid_at)
        known_instant = parse_instant(known_at)
        trusted = set(trust_domains)
        matches = [
            fact
            for fact in self._facts
            if fact.key == key
            and fact.valid_at(valid_instant)
            and fact.known_at(known_instant)
            and fact.accessible_to(principals)
            and (not trusted or fact.trust_domain in trusted)
        ]
        matches.sort(
            key=lambda fact: (
                parse_instant(fact.observed_at),
                parse_instant(fact.valid_from),
                fact.version_id,
            ),
            reverse=True,
        )
        values = tuple(sorted({fact.value for fact in matches}))
        return TemporalLookup(
            key=key,
            valid_at=valid_at,
            known_at=known_at,
            facts=tuple(matches),
            ambiguous=len(values) > 1,
            values=values,
        )

    def latest_known(
        self,
        key: str,
        known_at: str,
        principals: Iterable[str] = (),
    ) -> Optional[TemporalFact]:
        """Latest non-retracted version by observation time, regardless of valid time."""

        instant = parse_instant(known_at)
        matches = [
            fact
            for fact in self._facts
            if fact.key == key and fact.known_at(instant) and fact.accessible_to(principals)
        ]
        if not matches:
            return None
        return max(matches, key=lambda fact: (parse_instant(fact.observed_at), fact.version_id))

    def change_log(self, source_id: str = "") -> Tuple[TemporalFact, ...]:
        """Return versions in transaction-time order for replay and audit."""

        facts = [fact for fact in self._facts if not source_id or fact.source_id == source_id]
        facts.sort(key=lambda fact: (parse_instant(fact.observed_at), fact.version_id))
        return tuple(facts)


def age_seconds(observed_at: str, query_time: str) -> float:
    """Non-negative age of indexed evidence at query time."""

    delta = (parse_instant(query_time) - parse_instant(observed_at)).total_seconds()
    return max(0.0, delta)


def exponential_time_decay(
    score: float,
    observed_at: str,
    query_time: str,
    half_life_seconds: float,
) -> float:
    """Apply exponential recency decay without changing the sign of a score."""

    if half_life_seconds <= 0:
        raise ValueError("half_life_seconds must be positive")
    decay = math.exp(-math.log(2.0) * age_seconds(observed_at, query_time) / half_life_seconds)
    return score * decay


def freshness_sla_met(observed_at: str, query_time: str, maximum_age_seconds: float) -> bool:
    """Whether evidence age stays inside an explicit ingestion/query SLA."""

    if maximum_age_seconds < 0:
        raise ValueError("maximum_age_seconds must be non-negative")
    return age_seconds(observed_at, query_time) <= maximum_age_seconds


def stale_rate(
    observed_times: Sequence[str],
    query_time: str,
    maximum_age_seconds: float,
) -> float:
    """Fraction of returned evidence units violating the freshness SLA."""

    if not observed_times:
        return 0.0
    stale = sum(
        not freshness_sla_met(observed, query_time, maximum_age_seconds)
        for observed in observed_times
    )
    return stale / len(observed_times)


def cache_identity(
    query: str,
    corpus_snapshot: str,
    permission_fingerprint: str,
    valid_at: str,
    model_revision: str,
) -> Tuple[str, ...]:
    """Fields that must separate temporal and authorization-sensitive caches."""

    if not all((query, corpus_snapshot, permission_fingerprint, valid_at, model_revision)):
        raise ValueError("all cache identity fields must be non-empty")
    return (query, corpus_snapshot, permission_fingerprint, valid_at, model_revision)
