"""Memory Tools — what the agent carries between runs.

Both AUTOMATIC. Memory holds the agent's own notes, not client records: a wrong
one is a bad note that a human can correct, not a corrupted file.

Every write goes through ``assert_no_restricted``. Memory is the natural thing to
embed and later send to a hosted model for reasoning, so a CNIC written here
would be a CNIC leaving the installation the first time anyone asked a question
about that client (ADR-0002).
"""

from __future__ import annotations

from typing import Any

from app.memory.records import MemoryKind, MemoryRecord
from app.memory.repository import MemoryRepository
from app.runtime.metrics import METRICS
from app.security.redaction import RestrictedDataError, assert_no_restricted, redact_text
from app.tools.base import ExecutionPolicy, Tool, ToolResult


class RememberTool(Tool):
    """Write something down for a later run to find."""

    name = "remember"
    description = (
        "Record what was learned — a document filed, a decision made, something still missing. "
        "Identifiers are masked before storage."
    )
    policy = ExecutionPolicy.AUTOMATIC

    def __init__(self, memory: MemoryRepository) -> None:
        self._memory = memory

    def run(self, **kwargs: Any) -> ToolResult:
        content = (kwargs.get("content") or "").strip()

        if not content:
            return ToolResult.failure("A memory needs content.")

        try:
            kind = MemoryKind(kwargs.get("kind", MemoryKind.OBSERVATION))
        except ValueError:
            return ToolResult.failure(
                f"Unknown memory kind '{kwargs.get('kind')}'. "
                f"Expected one of: {', '.join(sorted(k.value for k in MemoryKind))}."
            )

        metadata = dict(kwargs.get("metadata") or {})

        try:
            # Checked before the record exists, not after: a record that has to
            # be deleted was still written down.
            assert_no_restricted(metadata, path="metadata")
        except RestrictedDataError as exc:
            return ToolResult.failure(str(exc))

        record = self._memory.remember(
            MemoryRecord(
                kind=kind,
                # Masked rather than rejected: prose is where an identifier turns
                # up by accident, and refusing the whole memory would lose the
                # note over a detail it did not need.
                content=redact_text(content),
                client_id=kwargs.get("client_id"),
                document_id=kwargs.get("document_id"),
                workflow_id=kwargs.get("workflow_id"),
                metadata=metadata,
                confidence=kwargs.get("confidence"),
            )
        )

        return ToolResult.success({"memory_id": record.id, "kind": record.kind.value})


class RecallTool(Tool):
    """Find what earlier runs learned about this client or question."""

    name = "recall"
    description = (
        "Retrieve relevant memories, most similar first. Optionally filtered by client or kind."
    )
    policy = ExecutionPolicy.AUTOMATIC

    def __init__(self, memory: MemoryRepository) -> None:
        self._memory = memory

    def run(self, **kwargs: Any) -> ToolResult:
        query = (kwargs.get("query") or "").strip()

        if not query:
            return ToolResult.failure("A query is required to recall anything.")

        kind = kwargs.get("kind")

        try:
            kind = MemoryKind(kind) if kind else None
        except ValueError:
            return ToolResult.failure(f"Unknown memory kind '{kind}'.")

        with METRICS.time("taxpilot_memory_recall_seconds", "Time spent recalling memories."):
            hits = self._memory.recall(
                query,
                limit=kwargs.get("limit", 5),
                client_id=kwargs.get("client_id"),
                kind=kind,
                min_similarity=kwargs.get("min_similarity", 0.0),
            )

        # Recall that returns nothing every time means memory is not earning its
        # place — a distinction invisible from timing alone.
        METRICS.counter("taxpilot_memory_recalls_total", "Memory recalls.",
                        outcome="hit" if hits else "empty")

        return ToolResult.success(
            {
                "memories": [
                    {
                        "id": record.id,
                        "kind": record.kind.value,
                        "content": record.content,
                        "client_id": record.client_id,
                        "similarity": round(score, 4),
                        "created_at": record.created_at.isoformat(),
                    }
                    for record, score in hits
                ],
                "count": len(hits),
            }
        )


class OpenGapsTool(Tool):
    """What is still outstanding — the basis for chasing a client."""

    name = "open_gaps"
    description = "List documents or information still missing, oldest first."
    policy = ExecutionPolicy.AUTOMATIC

    def __init__(self, memory: MemoryRepository) -> None:
        self._memory = memory

    def run(self, **kwargs: Any) -> ToolResult:
        gaps = self._memory.open_gaps(client_id=kwargs.get("client_id"))

        return ToolResult.success(
            {
                "gaps": [
                    {
                        "id": record.id,
                        "content": record.content,
                        "client_id": record.client_id,
                        "since": record.created_at.isoformat(),
                    }
                    for record in gaps
                ],
                "count": len(gaps),
            }
        )
