"""Storing and recalling memory.

A port, plus an in-memory implementation. Postgres and pgvector are the
production store (ADR-0001), but everything above this line is written against
the interface, so retrieval logic is testable without a database and a
deployment can run before one is provisioned.
"""

from __future__ import annotations

from typing import Protocol, runtime_checkable

from app.memory.embeddings import Embedder, HashingEmbedder, cosine_similarity
from app.memory.records import MemoryKind, MemoryRecord


@runtime_checkable
class MemoryRepository(Protocol):
    """What any memory store must provide."""

    def remember(self, record: MemoryRecord) -> MemoryRecord: ...

    def recall(
        self,
        query: str,
        *,
        limit: int = 5,
        client_id: int | None = None,
        kind: MemoryKind | None = None,
        min_similarity: float = 0.0,
    ) -> list[tuple[MemoryRecord, float]]: ...

    def open_gaps(self, client_id: int | None = None) -> list[MemoryRecord]: ...

    def get(self, record_id: str) -> MemoryRecord | None: ...

    def forget(self, record_id: str) -> bool: ...


class InMemoryRepository:
    """A repository held in process.

    Real enough to develop and test retrieval against, and honest about what it
    is: nothing survives a restart. Used by the test suite, and as the fallback
    for a deployment whose Postgres is not yet up — where losing memory on
    restart is far better than refusing to run.
    """

    def __init__(self, embedder: Embedder | None = None) -> None:
        self._embedder = embedder or HashingEmbedder()
        self._records: dict[str, MemoryRecord] = {}
        self._vectors: dict[str, list[float]] = {}

    def remember(self, record: MemoryRecord) -> MemoryRecord:
        self._records[record.id] = record
        self._vectors[record.id] = self._embedder.embed(record.content)

        return record

    def recall(
        self,
        query: str,
        *,
        limit: int = 5,
        client_id: int | None = None,
        kind: MemoryKind | None = None,
        min_similarity: float = 0.0,
    ) -> list[tuple[MemoryRecord, float]]:
        """Most relevant memories first, with their similarity.

        The score is returned rather than hidden so a caller can decide whether
        a weak match is worth acting on. A bare list of records invites treating
        the top hit as relevant when everything scored 0.02.

        Filters apply *before* ranking: a client filter that ranked first and
        filtered second would return fewer results than asked for whenever the
        best matches belonged to someone else.
        """
        candidates = [
            record
            for record in self._records.values()
            if (client_id is None or record.client_id == client_id)
            and (kind is None or record.kind is kind)
        ]

        query_vector = self._embedder.embed(query)

        scored = [
            (record, cosine_similarity(query_vector, self._vectors[record.id]))
            for record in candidates
        ]

        scored = [(r, s) for r, s in scored if s >= min_similarity]
        scored.sort(key=lambda pair: (-pair[1], pair[0].created_at))

        return scored[:limit]

    def open_gaps(self, client_id: int | None = None) -> list[MemoryRecord]:
        """What is still missing.

        Not a search: this is a definite question with a definite answer, and
        answering it by similarity would occasionally omit a gap because its
        wording did not match. Oldest first — the longest-outstanding request is
        the one most worth chasing.
        """
        gaps = [
            record
            for record in self._records.values()
            if record.is_open_gap and (client_id is None or record.client_id == client_id)
        ]

        return sorted(gaps, key=lambda r: r.created_at)

    def get(self, record_id: str) -> MemoryRecord | None:
        return self._records.get(record_id)

    def forget(self, record_id: str) -> bool:
        if record_id not in self._records:
            return False

        del self._records[record_id]
        self._vectors.pop(record_id, None)

        return True

    def __len__(self) -> int:
        return len(self._records)
