"""Conversation history store — Redis-backed, last N turns per session."""
from __future__ import annotations

import json
import logging
import time

logger = logging.getLogger("apex")

_MAX_TURNS = 10
_VERBATIM_TAIL = 4
_TTL_SECONDS = 3600


class ConversationStore:
    """Stores and retrieves conversation turns keyed by session token."""

    def __init__(self, redis_client, max_turns: int = _MAX_TURNS):
        self._redis = redis_client
        self._max_turns = max_turns

    # ── internal ───────────────────────────────────────────────────────────

    def _key(self, session_token: str) -> str:
        return f"conv:{session_token}"

    def _load(self, session_token: str) -> list[dict]:
        try:
            raw = self._redis.get(self._key(session_token))
            if raw:
                return json.loads(raw)
        except Exception as exc:
            logger.warning("[ConvStore] load error: %s", exc)
        return []

    def _save(self, session_token: str, turns: list[dict]) -> None:
        try:
            self._redis.set(
                self._key(session_token),
                json.dumps(turns, ensure_ascii=False),
                ex=_TTL_SECONDS,
            )
        except Exception as exc:
            logger.warning("[ConvStore] save error: %s", exc)

    # ── public API ─────────────────────────────────────────────────────────

    def add_turn(self, session_token: str, user_input: str, assistant_reply: str) -> None:
        """Append a turn; compress overflow into a rolling summary slot."""
        turns = self._load(session_token)
        turns.append({
            "ts": int(time.time()),
            "user": user_input[:500],
            "assistant": assistant_reply[:800],
        })
        if len(turns) > self._max_turns:
            to_compress = turns[: len(turns) - self._max_turns + 1]
            rest = turns[len(turns) - self._max_turns + 1 :]
            summary_lines = [
                "Q: " + t["user"][:80] + " -> A: " + t["assistant"][:120]
                for t in to_compress
                if not t.get("_summary")
            ]
            summary_entry = {
                "ts": to_compress[0]["ts"],
                "_summary": True,
                "user": "[Earlier context]",
                "assistant": " | ".join(summary_lines),
            }
            turns = [summary_entry] + rest
        self._save(session_token, turns)

    def get_history(self, session_token: str) -> list[dict]:
        """Return up to max_turns previous turns (oldest first)."""
        return self._load(session_token)

    def format_for_prompt(self, session_token: str) -> str:
        """Return history for LLM injection.

        Older turns (beyond the verbatim tail) are compressed into compact
        one-liners so the prompt stays concise while retaining context.
        """
        turns = self._load(session_token)
        if not turns:
            return ""
        verbatim = turns[-_VERBATIM_TAIL:]
        older = turns[: len(turns) - len(verbatim)]
        lines = []
        if older:
            compact = " | ".join(
                t["assistant"]
                if t.get("_summary")
                else "Q: " + t["user"][:80] + " -> A: " + t["assistant"][:120]
                for t in older
            )
            lines.append("[Earlier] " + compact)
        for t in verbatim:
            if t.get("_summary"):
                lines.append("[Earlier context] " + t["assistant"])
            else:
                lines.append("User: " + t["user"])
                lines.append("Apex: " + t["assistant"])
        return "\n".join(lines)

    def clear(self, session_token: str) -> None:
        try:
            self._redis.delete(self._key(session_token))
        except Exception as exc:
            logger.warning("[ConvStore] clear error: %s", exc)
