"""Auditable long-term memory lifecycle for conversational RAG agents."""

import math
from dataclasses import dataclass, replace
from typing import Iterable, List, Mapping, Optional, Sequence, Set, Tuple

from .temporal import parse_instant
from .text import content_terms, jaccard


@dataclass(frozen=True)
class MemoryRecord:
    """Externalized memory with provenance, permissions, and update semantics."""

    identifier: str
    text: str
    created_at: str
    source_turn: str
    kind: str = "episodic"
    importance: float = 0.5
    principals: Tuple[str, ...] = ()
    supersedes: Tuple[str, ...] = ()
    expires_at: str = ""
    deleted_at: str = ""
    last_accessed_at: str = ""
    access_count: int = 0

    def __post_init__(self) -> None:
        if not self.identifier or not self.text.strip() or not self.source_turn:
            raise ValueError("memory identifier, text, and source turn must be non-empty")
        if self.kind not in {"episodic", "semantic", "procedural", "profile", "summary"}:
            raise ValueError("unsupported memory kind")
        if not 0.0 <= self.importance <= 1.0:
            raise ValueError("importance must be between zero and one")
        if self.access_count < 0:
            raise ValueError("access_count must be non-negative")
        parse_instant(self.created_at)
        if self.expires_at:
            parse_instant(self.expires_at)
        if self.deleted_at:
            parse_instant(self.deleted_at)

    def accessible_to(self, principals: Iterable[str]) -> bool:
        allowed = set(self.principals)
        return not allowed or bool(allowed & set(principals))

    def active_at(self, instant: str) -> bool:
        point = parse_instant(instant)
        if parse_instant(self.created_at) > point:
            return False
        if self.deleted_at and parse_instant(self.deleted_at) <= point:
            return False
        if self.expires_at and parse_instant(self.expires_at) <= point:
            return False
        return True


@dataclass(frozen=True)
class MemoryHit:
    """Retrieved memory plus decomposed relevance features."""

    record: MemoryRecord
    score: float
    lexical: float
    recency: float
    importance: float


@dataclass(frozen=True)
class WriteDecision:
    """Memory-write policy decision with inspectable reasons."""

    write: bool
    novelty: float
    reasons: Tuple[str, ...]


class MemoryStore:
    """Append/update/delete store that keeps memory lifecycle observable."""

    def __init__(self, records: Sequence[MemoryRecord] = ()) -> None:
        self._records: List[MemoryRecord] = []
        self._ids: Set[str] = set()
        for record in records:
            self.append(record)

    @property
    def records(self) -> Tuple[MemoryRecord, ...]:
        return tuple(self._records)

    def append(self, record: MemoryRecord) -> None:
        if record.identifier in self._ids:
            raise ValueError("duplicate memory identifier: " + record.identifier)
        missing = set(record.supersedes) - self._ids
        if missing:
            raise ValueError("memory supersedes unknown records: " + ", ".join(sorted(missing)))
        self._records.append(record)
        self._ids.add(record.identifier)

    def delete(self, identifier: str, deleted_at: str) -> None:
        """Logical deletion retained in the audit log."""

        parse_instant(deleted_at)
        for index, record in enumerate(self._records):
            if record.identifier == identifier:
                if record.deleted_at:
                    raise ValueError("memory is already deleted")
                self._records[index] = replace(record, deleted_at=deleted_at)
                return
        raise KeyError(identifier)

    def retrieve(
        self,
        query: str,
        at: str,
        principals: Iterable[str] = (),
        k: int = 5,
        recency_half_life_days: float = 30.0,
    ) -> Tuple[MemoryHit, ...]:
        """Retrieve active, authorized, non-superseded memories."""

        if k <= 0:
            return ()
        if recency_half_life_days <= 0:
            raise ValueError("recency half-life must be positive")
        point = parse_instant(at)
        active = [
            record
            for record in self._records
            if record.active_at(at) and record.accessible_to(principals)
        ]
        superseded = {identifier for record in active for identifier in record.supersedes}
        query_terms = set(content_terms(query))
        hits = []
        for record in active:
            if record.identifier in superseded:
                continue
            terms = set(content_terms(record.text))
            lexical = len(query_terms & terms) / max(1, len(query_terms))
            age_days = max(0.0, (point - parse_instant(record.created_at)).total_seconds() / 86400)
            recency = math.exp(-math.log(2.0) * age_days / recency_half_life_days)
            score = 0.65 * lexical + 0.20 * record.importance + 0.15 * recency
            if lexical > 0 or record.importance >= 0.8:
                hits.append(MemoryHit(record, score, lexical, recency, record.importance))
        hits.sort(key=lambda hit: (-hit.score, hit.record.identifier))
        return tuple(hits[:k])

    def mark_accessed(self, identifiers: Iterable[str], at: str) -> None:
        """Update access metadata without altering memory creation time."""

        parse_instant(at)
        requested = set(identifiers)
        found: Set[str] = set()
        for index, record in enumerate(self._records):
            if record.identifier in requested:
                self._records[index] = replace(
                    record,
                    last_accessed_at=at,
                    access_count=record.access_count + 1,
                )
                found.add(record.identifier)
        missing = requested - found
        if missing:
            raise KeyError(", ".join(sorted(missing)))

    def purgeable(self, at: str) -> Tuple[str, ...]:
        """IDs whose retention period expired or whose logical deletion is effective."""

        point = parse_instant(at)
        purge = []
        for record in self._records:
            expired = bool(record.expires_at and parse_instant(record.expires_at) <= point)
            deleted = bool(record.deleted_at and parse_instant(record.deleted_at) <= point)
            if expired or deleted:
                purge.append(record.identifier)
        return tuple(sorted(purge))


def write_decision(
    candidate_text: str,
    existing: Sequence[MemoryRecord],
    importance: float,
    durable_signal: bool = False,
    novelty_threshold: float = 0.35,
) -> WriteDecision:
    """Decide whether a turn deserves durable memory.

    A write is justified by novelty plus importance, or by an explicit durable
    signal such as a user preference/correction.  This prevents storing every
    turn while keeping policy decisions explainable.
    """

    if not 0.0 <= importance <= 1.0 or not 0.0 <= novelty_threshold <= 1.0:
        raise ValueError("importance and novelty threshold must be between zero and one")
    terms = set(content_terms(candidate_text))
    similarity = max(
        (jaccard(terms, set(content_terms(record.text))) for record in existing),
        default=0.0,
    )
    novelty = 1.0 - similarity
    reasons = []
    if novelty >= novelty_threshold:
        reasons.append("novel")
    if importance >= 0.7:
        reasons.append("important")
    if durable_signal:
        reasons.append("explicit-durable-signal")
    write = durable_signal or (novelty >= novelty_threshold and importance >= 0.5)
    if not write:
        reasons.append("write-threshold-not-met")
    return WriteDecision(write, novelty, tuple(reasons))


def consolidation_groups(
    records: Sequence[MemoryRecord], threshold: float = 0.6
) -> Tuple[Tuple[str, ...], ...]:
    """Connected components of semantically overlapping memories.

    This returns candidates for a summarizer; it deliberately does not generate
    a lossy summary or delete originals.  Human/audited consolidation can then
    create a new record whose ``supersedes`` field preserves lineage.
    """

    if not 0.0 <= threshold <= 1.0:
        raise ValueError("threshold must be between zero and one")
    terms: Mapping[str, Set[str]] = {
        record.identifier: set(content_terms(record.text)) for record in records
    }
    adjacency = {record.identifier: set() for record in records}
    for left_index, left in enumerate(records):
        for right in records[left_index + 1 :]:
            if jaccard(terms[left.identifier], terms[right.identifier]) >= threshold:
                adjacency[left.identifier].add(right.identifier)
                adjacency[right.identifier].add(left.identifier)
    groups = []
    seen: Set[str] = set()
    for identifier in sorted(adjacency):
        if identifier in seen:
            continue
        stack = [identifier]
        component = []
        while stack:
            current = stack.pop()
            if current in seen:
                continue
            seen.add(current)
            component.append(current)
            stack.extend(sorted(adjacency[current] - seen, reverse=True))
        if len(component) > 1:
            groups.append(tuple(sorted(component)))
    return tuple(groups)
